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

PyTorch转ONNX后OnnxRuntime报Trilu/Where未实现错误的解决问询

解决PyTorch模型转ONNX后OnnxRuntime的Trilu/Where未实现问题

1. 升级ONNX opset与Runtime版本

导出ONNX时指定opset 16及以上,同时确保OnnxRuntime版本在1.13+,因为Trilu节点在高opset中才被OnnxRuntime完整支持:

torch.onnx.export(
    model, 
    input_sample, 
    "model.onnx", 
    opset_version=16, 
    export_params=True,
    do_constant_folding=True
)

do_constant_folding参数会让PyTorch提前优化常量节点,减少动态生成逻辑带来的兼容性问题。

2. 重写tril逻辑为ONNX友好的实现

如果高opset仍有问题,避免使用torch.tril或手动Where实现,改用Range/GreaterOrEqual生成掩码后做元素乘法:

def custom_tril(x, diagonal=0):
    batch_size, seq_len, _ = x.shape
    rows = torch.arange(seq_len, device=x.device).unsqueeze(1)
    cols = torch.arange(seq_len, device=x.device).unsqueeze(0)
    # 生成下三角掩码:行索引 >= 列索引 - diagonal
    mask = (rows >= cols - diagonal).float()
    mask = mask.unsqueeze(0).repeat(batch_size, 1, 1)
    return x * mask

替换模型中原有的torch.tril调用,导出的ONNX会生成OnnxRuntime原生支持的Range、GreaterOrEqual和Mul节点,避开Trilu/Where的兼容性问题。

3. 后处理已有ONNX模型

如果已经导出了问题模型,可以用ONNX工具修改节点:

import onnx
from onnx import helper

model = onnx.load("model.onnx")
# 遍历节点找到Trilu节点,替换为掩码+Mul的组合
# 示例逻辑(需根据实际模型结构调整):
for node in model.graph.node:
    if node.op_type == "Trilu":
        # 获取输入输出与diagonal参数
        input_tensor = node.input[0]
        output_tensor = node.output[0]
        diagonal = node.attribute[0].i if node.attribute else 0
        
        # 根据模型张量形状生成掩码节点(Range、GreaterOrEqual、Cast等)
        # 替换原Trilu节点后删除旧节点
        
onnx.save(model, "fixed_model.onnx")

环境核查

  • PyTorch版本建议1.12+,避免旧版本的ONNX导出bug
  • OnnxRuntime版本必须匹配opset支持范围,比如opset16需要Runtime1.13+

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 10:20:39