已开启
FlexAttention 动态 mask 捕获符号 shape 时编译失败 #4425
Xuan Peng创建于  3 天前
Xuan Peng
3 天前 创建

环境信息

  • 目标分支:master(PyTorch 2.13 适配)
  • 目标硬件:Ascend 950
  • 参考验证硬件:Ascend 910B4(A2-140,仅用于编译和功能参考)

问题描述

FlexAttention 的 mask_modscore_mod 捕获动态张量 shape 时,捕获张量可能只贡献符号尺寸,并不作为 Triton kernel 的位置参数参与实际读取。

当前 lowering 将原始捕获列表继续传给 autotune benchmark,可能导致 benchmark 输入数量、顺序与模板实际 input_nodes 不一致。在 A2 参考环境中,该问题表现为:

launcher() got multiple values for argument 'stream'

同时,编译阶段读取并缓存 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)

预期行为

  • 编译期间不读取或缓存数据相关的 eager BlockMask 内容。
  • autotune benchmark 输入与模板实际 kernel 参数保持一致。
  • 仅用于 shape 的符号捕获作为图依赖保留,不占用 tensor 位置参数。
  • 动态窗口大小复用同一张编译图,并保持结果正确。

建议方案与验收标准

  1. 根据候选模板的 input_nodes 对齐 autotune tensor 输入,并排除 SymPy 标量。
  2. 在模板选择结果上记录实际执行子图的输入、输出,保留自由符号依赖。
  3. 移除 eager BlockMask 元数据分析与缓存,复用 master 已有的运行时精确稀疏容量逻辑。
  4. 增加动态 mask_mod shape 的回归测试。
  5. 不为 A2/910B 单独收缩生产 tile;A2 的 UB 或运行限制仅作为参考,最终策略面向 Ascend 950。
likedislike
TorchNPU-BotTorchNPU-Bot成员
3 天前 添加了label:triage-review
TorchNPU-Bot
TorchNPU-Bot成员
3 天前 评论:

issue待分派,添加triage-review标签

likedislike
XXuan Peng
3 天前 关联了pull request:fix(flexattention): avoid metadata guards and preserve subgraph symbol uses
TorchNPU-BotTorchNPU-Bot成员
3 天前 添加了label:bot-triaged;删除了label:triage-review
TorchNPU-Bot
TorchNPU-Bot成员
3 天前 评论:

检测到当前 issue 已关联 PR,自动添加标签:bot-triaged

likedislike