训练前如何修改PyTorch ResNet50输入形状,解决TFLite转译Flutter报错
解决PyTorch ResNet50输入形状调整及TFLite转换报错问题
一、训练前让模型接受(224,224,3)格式输入
PyTorch默认采用通道在前(CHW,即(3,224,224))的输入格式,ResNet50的核心卷积层也是基于此设计的。要让模型兼容通道在后(HWC,即(224,224,3))的输入,无需改动ResNet50的核心结构,只需在模型的前向传播中添加维度转置操作即可:
import torch import torchvision.models as models class HWCResNet50(torch.nn.Module): def __init__(self, pretrained=False): super().__init__() self.resnet = models.resnet50(pretrained=pretrained) def forward(self, x): # 输入x形状为 (batch_size, 224, 224, 3) x = x.permute(0, 3, 1, 2) # 转置为PyTorch所需的 (batch_size, 3, 224, 224) return self.resnet(x) # 实例化并测试模型 model = HWCResNet50(pretrained=True) test_input = torch.randn(1, 224, 224, 3) output = model(test_input) print(output.shape) # 输出应为 (1, 1000),符合预期
训练时直接传入HWC格式的图像张量即可,无需额外预处理转置。
二、转换为TFLite时解决输入形状不匹配问题
训练完成后,需确保导出的TFLite模型输入形状为(1,224,224,3),可通过以下两种方法实现:
方法1:ONNX中转法
先将PyTorch模型导出为ONNX格式,明确指定HWC输入形状,再转换为TFLite:
- 导出ONNX模型
model.eval() dummy_input = torch.randn(1, 224, 224, 3) # 用HWC格式的dummy输入 torch.onnx.export( model, dummy_input, "resnet50_hwc.onnx", input_names=["input"], output_names=["output"], dynamic_axes={"input": {0: "batch_size"}, "output": {0: "batch_size"}} # 支持动态batch大小 )
- 转换为TFLite
import tensorflow as tf converter = tf.lite.TFLiteConverter.from_onnx("resnet50_hwc.onnx") converter.optimizations = [tf.lite.Optimize.DEFAULT] # 可选:启用优化 tflite_model = converter.convert() # 保存模型 with open("resnet50_hwc.tflite", "wb") as f: f.write(tflite_model)
方法2:直接转换为TensorFlow函数再转TFLite
无需ONNX中转,直接将PyTorch模型转为TensorFlow兼容格式:
import torch import tensorflow as tf from torch.utils.tensorflow import convert_to_tf_tensor model.eval() dummy_input = torch.randn(1, 224, 224, 3) # 将PyTorch模型转为TensorFlow追踪函数 tf_func = tf.function(lambda x: convert_to_tf_tensor(model(torch.from_numpy(x.numpy())))) concrete_func = tf_func.get_concrete_function(tf.TensorSpec([1, 224, 224, 3], tf.float32)) # 转换为TFLite converter = tf.lite.TFLiteConverter.from_concrete_functions([concrete_func]) tflite_model = converter.convert() # 保存模型 with open("resnet50_hwc.tflite", "wb") as f: f.write(tflite_model)
验证TFLite输入形状
转换完成后,可验证输入张量是否符合要求:
interpreter = tf.lite.Interpreter(model_path="resnet50_hwc.tflite") interpreter.allocate_tensors() input_details = interpreter.get_input_details() print(input_details[0]['shape']) # 应输出 [1, 224, 224, 3]
内容的提问来源于stack exchange,提问作者borlor andurar
相关产品推荐
相关产品推荐

