You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

训练前如何修改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:

  1. 导出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大小
)
  1. 转换为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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.28 01:05:09