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

Matlab预训练ResNet50转ONNX后PyTorch推理报错及可行性咨询

问题修复与解答

报错修复

你遇到的sum()参数错误,主要有两个可能的原因,对应修复方案如下:

1. 模型输出格式问题

Matlab导出的ONNX模型可能包含多个输出节点,通过onnx2pytorch转换后,模型的forward方法会返回元组而非单个张量。你需要取元组的第一个元素作为模型的logits输出:

# 原代码
output = model(test_inputs)
# 修改为
output = model(test_inputs)[0]

2. 正确数计算的张量操作问题

原代码中torch.sum(pred == test_labels.data)的写法在部分PyTorch版本或张量维度不匹配时可能触发参数错误,同时.data属性已被PyTorch官方不推荐使用。建议修改为直接求和并转成Python数值:

# 原代码
testing_correct += torch.sum(pred == test_labels.data)
# 修改为
testing_correct += (pred == test_labels).sum().item()

额外的准确率计算错误

你的测试准确率计算逻辑有误:len(test_loader)是测试集的批次数量,不是总样本数,会导致准确率被错误放大。需要改为基于总样本数计算:

# 在循环外初始化总样本数
total_test_samples = 0

# 在批次循环内累加样本数
total_test_samples += test_inputs.size(0)

# 计算准确率
test_acc = 100 * testing_correct / total_test_samples

ONNX模型在PyTorch中推理的可行性

完全可行。onnx2pytorch这类工具可以将标准ONNX模型转换为PyTorch可执行的nn.Module实例,满足推理需求。需要注意以下几点:

  • 算子兼容性:Matlab预训练的ResNet50属于常见模型,所用算子基本都能被PyTorch兼容,转换不会有问题。
  • 输入预处理一致性:Matlab的ResNet50预训练时的输入归一化参数(均值、标准差)可能和PyTorch官方预训练模型不同,需要确保测试时的输入预处理和Matlab训练时一致,否则会影响精度。
  • 推理模式设置:必须调用model.eval()关闭BatchNorm、Dropout等层的训练行为,否则推理结果会异常。

内容的提问来源于stack exchange,提问作者Jacky

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.22 05:24:17