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

如何让torch.ones在symbolic_trace的追踪Wrapper模型中正常工作?

解决torch.fx.symbolic_trace追踪时torch.ones创建失败的问题

问题根源

torch.fx符号追踪过程中,input_ids.size(0)、input_ids.size(1)返回的是符号值(SymbolicInt),而非普通Python整数。直接将这些符号值作为tuple传入torch.ones()时,FX追踪机制无法正确解析动态形状的构造逻辑,从而触发索引相关错误。

可行解决方案

1. 使用torch.ones_like()替代手动构造形状

直接基于输入张量的形状生成全1掩码,完全避免手动提取维度值的操作,FX能直接追踪到形状依赖关系:

attention_mask = torch.ones_like(input_ids, dtype=torch.int64)

2. 手动构造形状时用torch.Size包装维度值

如果需要自定义形状(比如与输入维度不完全一致),用torch.Size包装符号维度值,让FX正确识别形状的符号化构造逻辑:

batch_size = input_ids.size(0)
seq_len = input_ids.size(1)
attention_mask = torch.ones(torch.Size((batch_size, seq_len)), dtype=torch.int64)

修改后的完整WrappedModel代码

import logging
import torch

class WrappedModel(torch.nn.Module):
    def __init__(self, model):
        super().__init__()
        self.model = model

    def forward(self, input_ids):
        try:
            # 推荐方案:用ones_like直接生成对应形状的mask
            attention_mask = torch.ones_like(input_ids, dtype=torch.int64)
            
            # 备选方案:手动构造形状时用torch.Size包装
            # batch_size = input_ids.size(0)
            # seq_len = input_ids.size(1)
            # attention_mask = torch.ones(torch.Size((batch_size, seq_len)), dtype=torch.int64)

            output = self.model(input_ids=input_ids, attention_mask=attention_mask)
        except Exception as e:
            logging.warning(f"TRACE ERROR inside wrapped forward: {e}")
            return torch.zeros(1, 1)

        if hasattr(output, "last_hidden_state"):
            return output.last_hidden_state
        elif hasattr(output, "logits"):
            return output.logits
        return output

补充说明

  • 之前尝试的硬编码形状虽能绕过错误,但失去了动态适配输入维度的能力;input_ids.size()和input_ids.shape本质都是返回符号值,替换后无法解决根本问题。
  • 确保追踪时的虚拟输入符合模型要求(如torch.randint(0, 1000, (1, 10))是合理的),同时模型本身没有其他阻碍FX追踪的操作(比如未被支持的自定义算子、动态控制流)。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 07:57:13