TensorFlow转TFLite输入形状异常:[1,320,320,3]变[1,1,1,3]求助
解决TensorFlow预训练模型转TFLite后输入形状异常问题
当转换SSD MobileNet v2这类预训练目标检测模型时,若未显式指定输入形状,TFLite转换器可能会将动态维度默认设为[1,1,1,3],导致与原模型输入形状不符。以下是针对性的解决方案:
1. 显式指定输入形状重新转换
修改转换代码,强制固定输入形状为原模型的[1,320,320,3]:
import tensorflow as tf # 加载SavedModel并获取输入张量名称 saved_model = tf.saved_model.load('/content/drive/MyDrive/ssd_mobilenet_v2_320x320_coco17_tpu-8/saved_model') infer_func = saved_model.signatures['serving_default'] input_tensor_name = infer_func.inputs[0].name # 初始化转换器并配置参数 converter = tf.lite.TFLiteConverter.from_saved_model('/content/drive/MyDrive/ssd_mobilenet_v2_320x320_coco17_tpu-8/saved_model') # 启用必要的操作支持,避免转换报错 converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS, tf.lite.OpsSet.SELECT_TF_OPS] converter.experimental_enable_resource_variables = True # 强制指定输入形状 converter.input_shapes = {input_tensor_name: [1, 320, 320, 3]} # 执行转换并保存模型 tflite_model = converter.convert() with open('model.tflite', 'wb') as f: f.write(tflite_model)
2. 验证转换结果
转换完成后,可通过以下代码确认输入形状是否正确:
import tensorflow as tf interpreter = tf.lite.Interpreter(model_path='model.tflite') interpreter.allocate_tensors() input_details = interpreter.get_input_details() print(f"转换后输入形状: {input_details[0]['shape']}") print(f"输入张量名称: {input_details[0]['name']}")
正常情况下会输出转换后输入形状: [ 1 320 320 3]。
关键说明
- 部分预训练目标检测模型的SavedModel默认使用动态输入维度,必须显式固定形状才能得到符合预期的TFLite模型
- 若模型存在多个输入,需为每个输入张量分别指定对应形状
- 配置
target_spec.supported_ops时包含SELECT_TF_OPS,可兼容部分TFLite原生不支持的TensorFlow操作
内容的提问来源于stack exchange,提问作者Prashanth
相关产品推荐
相关产品推荐

