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

PyTorch预训练.pth模型转CoreML报错:无法获取Python类对象名称

解决PyTorch转CoreML时的RuntimeError问题

问题根源

你直接用torch.load加载的.pth文件是模型的权重参数字典,而非完整的模型实例。torch.jit.trace需要接收一个继承自nn.Module的模型对象,而非权重字典,因此会触发"无法获取类名"的错误。

解决步骤

  1. 导入A2J的模型结构
    从A2J项目中找到对应模型的定义代码(比如项目里side.py中的SideModel类),将这段类定义复制到你的脚本中,确保代码能正确引用该模型类。

  2. 实例化模型并加载权重
    先创建模型的实例对象,再通过load_state_dict方法加载.pth文件中的权重,而非直接用torch.load加载整个模型。

  3. 设置模型为评估模式
    转换前必须将模型切换到eval()模式,避免训练时的特殊层(如Dropout、BatchNorm)改变推理行为,确保追踪结果准确。

  4. 重新执行模型追踪与转换
    用实例化后的模型执行torch.jit.trace,再转换为CoreML模型。

完整代码示例

import coremltools as ct
import torch
import torch.nn as nn

# 1. 复制A2J项目中SideModel的完整类定义到此处
class SideModel(nn.Module):
    def __init__(self):
        super(SideModel, self).__init__()
        # 此处替换为A2J项目中SideModel的真实层结构
        self.base = nn.Sequential(
            nn.Conv2d(3, 64, kernel_size=3, padding=1),
            nn.ReLU(inplace=True),
            nn.MaxPool2d(kernel_size=2, stride=2)
            # ... 其余层定义需和原项目完全一致
        )
        # 原模型的其他分支/层定义

    def forward(self, x):
        # 此处替换为A2J项目中SideModel的真实前向传播逻辑
        x = self.base(x)
        # ... 其余前向传播步骤
        return x

# 2. 实例化模型并加载权重
model = SideModel()
state_dict = torch.load('/Users/sarojraut/Downloads/side.pth', map_location=torch.device('cpu'))
model.load_state_dict(state_dict)

# 3. 切换到评估模式
model.eval()

# 4. 追踪模型并转换为CoreML
example_input = torch.rand(1, 3, 224, 224)
traced_model = torch.jit.trace(model, example_input)

# 转换时可根据需求配置输入类型(比如设置为图像输入并做归一化)
coreml_model = ct.convert(
    traced_model,
    inputs=[ct.ImageType(name="input", shape=example_input.shape, scale=1/255.0)]
)

# 保存CoreML模型
coreml_model.save('A2J_Side.mlmodel')

注意事项

  • 必须保证模型结构和预训练权重完全匹配,否则会出现权重加载失败的问题。
  • 如果原模型包含自定义层或特殊操作,需要确保这些操作能被TorchScript兼容,或者转换为CoreML支持的操作。
  • 输入张量的形状要和模型训练时的输入一致,A2J默认输入为(1,3,224,224),若有差异需同步调整。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 02:05:18