如何让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
相关产品推荐
相关产品推荐

