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

PyTorch转ONNX后src_key_padding_mask未降低推理耗时问题

PyTorch Transformer转ONNX后掩码未带来推理加速问题

问题背景

我在PyTorch中训练了两个Transformer模型,导出为ONNX格式做推理,二者唯一区别是其中一个额外指定了src_key_padding_mask输入:

导出代码

  • 无掩码模型:
torch.onnx.export(model.cpu(), X.view(1,MAXSEQLEN), "model.onnx", verbose=False, input_names=['model_inputs'], output_names=['model_outputs'])
  • 带掩码输入模型:
torch.onnx.export(model.cpu(), (X.view(1,MAXSEQLEN), X_pad_mask.view(1,MAXSEQLEN)), "model-with-padmask.onnx", verbose=False, input_names=['model_inputs', 'padding_mask'], output_names=['model_outputs'])

核心问题

使用填充掩码时,未观察到两个模型的推理耗时有差异。我测试了从几乎全填充到无填充的各种输入,预期填充/掩码量增加时,因更多计算被屏蔽,推理会更快,但实际并非如此。

导出时的警告信息

C:\Users<user>\Miniconda3\envs\py10\lib\site-packages\torch\onnx\utils.py:620: UserWarning: ONNX Preprocess - Removing mutation from node aten::masked_fill_ on block input: 'src_key_padding_mask'. This changes graph semantics. (Triggered internally at ..\torch\csrc\jit\passes\onnx\remove_inplace_ops_for_onnx.cpp:355.) _C._jit_pass_onnx_remove_inplace_ops_for_onnx(graph, module)

后续测试及结果

遍历不同最大序列长度,观察三种场景的推理耗时:

  • (i) 输入x和全为False的x_mask(等效无掩码)
  • (ii) 输入x和半为True的x_mask(屏蔽x后半部分)
  • (iii) 仅输入x(src_padding_mask=None)

结果显示:

  • (i)和(ii)耗时相近;
  • (iii)在短序列时耗时与前两者相近,但长序列时反而更快,这和预期(无掩码会增加长序列耗时)完全相反。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.09 05:50:16