已开启
FlexAttention 动态 mask 捕获符号 shape 时编译失败 #4425
Xuan Peng创建于 3 天前
3 天前 添加了label:triage-review
TorchNPU-Bot
3 天前 评论:
3 天前 评论:
issue待分派,添加triage-review标签


3 天前 添加了label:bot-triaged;删除了label:triage-review
TorchNPU-Bot
3 天前 评论:
3 天前 评论:
检测到当前 issue 已关联 PR,自动添加标签:bot-triaged


环境信息
master(PyTorch 2.13 适配)问题描述
FlexAttention 的
mask_mod或score_mod捕获动态张量 shape 时,捕获张量可能只贡献符号尺寸,并不作为 Triton kernel 的位置参数参与实际读取。当前 lowering 将原始捕获列表继续传给 autotune benchmark,可能导致 benchmark 输入数量、顺序与模板实际
input_nodes不一致。在 A2 参考环境中,该问题表现为:同时,编译阶段读取并缓存 eager
BlockMask内容会引入数据相关的元数据特化;模板选择结果若未记录实际执行子图的输入、输出,也可能丢失动态 shape 的自由符号依赖。复现方式
使用动态窗口大小构造
BlockMask,并让mask_mod捕获window_source.shape[0]:def window_mask(_b, _h, q_idx, kv_idx): return (q_idx - kv_idx).abs() <= window_source.shape[0] compiled = torch.compile(fn, backend="inductor", dynamic=True, fullgraph=True) for window_size in (32, 48): window_source = torch.randn(window_size, device="npu") torch._dynamo.mark_dynamic(window_source, 0) block_mask = create_block_mask(window_mask, B=1, H=1, Q_LEN=128, KV_LEN=128, device="npu") compiled(q, k, v, block_mask, window_source)预期行为
BlockMask内容。建议方案与验收标准
input_nodes对齐 autotune tensor 输入,并排除 SymPy 标量。mask_modshape 的回归测试。