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

PyTorch Lightning训练的HuggingFace多标签分类模型导出ONNX失败求助

问题解决方法

报错根因

你的模型forward方法默认返回两个值,第一个是损失值loss,第二个是分类结果output。导出ONNX时你没有传入labels参数,所以loss被固定为常量0,和输入张量没有数据依赖关系,Torch的ONNX追踪器识别到无依赖的常量输出就会触发该报错。

具体解决步骤

  • 调整模型返回逻辑:ONNX导出用于推理场景,不需要计算损失,仅保留预测输出即可。
  • 修正输入配置:你的模型有两个输入input_ids和attention_mask,原有配置只写了一个输入名,和实际入参数量不匹配。
  • 统一设备:确保模型参数和输入张量在同一个设备上(全CPU或全GPU)。
  • 升级opset版本:opset10对Transformer类结构的支持不完善,建议升级到opset13以上适配BERT结构。

修改后的导出代码示例

方案1:重写forward方法(兼容所有PyTorch版本,推荐)

# 加载训练好的checkpoint权重(如果未加载先执行这步)
# model = SRTagger.load_from_checkpoint("你的checkpoint路径.ckpt", n_classes=100)

# 切换模型到推理模式
model.eval()
# 统一模型和输入的设备
device = next(model.parameters()).device
input_ids = sample_batch["input_ids"].to(device)
attention_mask = sample_batch["attention_mask"].to(device)

# 临时重写forward方法,仅返回推理需要的预测结果
original_forward = model.forward
def inference_forward(input_ids, attention_mask):
    _, output = original_forward(input_ids, attention_mask)
    return output
model.forward = inference_forward

# 执行导出
torch.onnx.export(
    model,
    (input_ids, attention_mask),
    "model_torch_export.onnx",
    export_params=True,
    opset_version=13,
    do_constant_folding=True,
    # 两个输入对应两个名称
    input_names = ['input_ids', 'attention_mask'],
    output_names = ['output'],
    dynamic_axes={
        'input_ids' : {0 : 'batch_size'},
        'attention_mask' : {0 : 'batch_size'},
        'output' : {0 : 'batch_size'}
    }
)

方案2:用output_indices参数过滤输出(PyTorch 1.8+支持)

如果你的PyTorch版本支持output_indices参数,可以不用重写forward,直接指定返回第二个输出即可:

torch.onnx.export(
    model,
    (input_ids, attention_mask),
    "model_torch_export.onnx",
    export_params=True,
    opset_version=13,
    do_constant_folding=True,
    input_names = ['input_ids', 'attention_mask'],
    output_names = ['output'],
    dynamic_axes={
        'input_ids' : {0 : 'batch_size'},
        'attention_mask' : {0 : 'batch_size'},
        'output' : {0 : 'batch_size'}
    },
    # 仅返回索引为1的输出,跳过无依赖的loss常量
    output_indices=[1]
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 02:09:02