如何正确向TensorFlow Lite模型传入字符串以实现推理?
解决方案:TFLite字符串输入维度不匹配问题
问题根源
从输入详情可以明确:
- 模型输入的静态形状为
[1],但shape_signature是[-1],说明支持动态批量输入 - 输入类型要求为
numpy.bytes_,需传入numpy字节数组,而非TensorFlow原生张量 - 报错核心是输入维度与模型期望不匹配:要么批量数不符,要么单个样本未包装成1维数组
修复步骤与可复现代码
1. 单个样本推理
将单个字符串包装成形状为(1,)的numpy字节数组即可:
import tensorflow as tf import numpy as np # 加载TFLite模型 interpreter = tf.lite.Interpreter(model_path="./final_model.tflite") interpreter.allocate_tensors() input_details = interpreter.get_input_details() output_details = interpreter.get_output_details() # 单个输入样本:转成(1,)的numpy字节数组 test_sample = "你的测试文本内容" input_data = np.array([test_sample.encode('utf-8')], dtype=np.bytes_) # 设置张量并执行推理 interpreter.set_tensor(input_details[0]['index'], input_data) interpreter.invoke() # 获取输出结果 output = interpreter.get_tensor(output_details[0]['index']) print("单个样本输出:", output)
2. 批量样本推理
利用模型支持动态批量的特性,先调整输入张量形状,再重新分配内存:
import tensorflow as tf import numpy as np # 加载TFLite模型 interpreter = tf.lite.Interpreter(model_path="./final_model.tflite") # 调整输入形状为目标批量大小(示例为2) batch_size = 2 input_details = interpreter.get_input_details() interpreter.resize_tensor_input(input_details[0]['index'], (batch_size,)) interpreter.allocate_tensors() # 调整形状后必须重新分配内存 output_details = interpreter.get_output_details() # 批量输入样本:转成(batch_size,)的numpy字节数组 test_samples = ["测试文本1", "测试文本2"] input_data = np.array([s.encode('utf-8') for s in test_samples], dtype=np.bytes_) # 设置张量并执行推理 interpreter.set_tensor(input_details[0]['index'], input_data) interpreter.invoke() # 获取输出结果 output = interpreter.get_tensor(output_details[0]['index']) print("批量样本输出:", output)
关键注意事项
- 输入类型严格匹配:模型要求
numpy.bytes_类型,需将字符串编码为字节(encode('utf-8'))并包装成numpy数组 - 维度严格对应:单个样本需为
(1,)形状,批量样本为(N,)(N为批量数),不能直接传入0维字符串或未转换的TensorFlow张量 - 动态批量需先调整形状:处理多个样本时,必须先调用
resize_tensor_input修改输入形状,再重新执行allocate_tensors
内容的提问来源于stack exchange,提问作者Bob
相关产品推荐
相关产品推荐

