已合并
flexattention: port flex attention overflow fixes to master #45274
a_knight创建于 8月25日
flexattention: port flex attention overflow fixes to master #45274
已合并
a_knight创建于 8月25日
a_knight
8月25日

【合入来源】

如有社区issue,请关联issue链接
请勿携带内部流程信息(需求链接、问题单、内部issue等)
https://gitcode.com/Ascend/pytorch/issues/4254
https://gitcode.com/Ascend/pytorch/issues/4341

【修改方案】

请描述修改内容的具体实现,涉及哪些组件之间进行交互,可以用1、2、3、...进行罗列
如果是需求或者重构类的PR,需要补充详细设计文档(说明上下游组件关系、时序图、类图、DFX能力等内容)

【资料变更】

请确认是否涉及资料变更。如涉及,需要在PR中体现,并简要说明修改内容。如不涉及,需填写“不涉及”

【接口变更】

请确认是否涉及跨代码仓或者客户面可见的接口变更。如涉及,需要详细说明接口以及对应的变更内容,同时需要在资料中体现。如不涉及,需填写“不涉及”

【功能验证】

说明测试场景,测试方法。如果本次测试方式与常规单元测试不同,请详细说明您的测试步骤
新增/变更内容是否已新增/适配UT测试用例看护,并补充测试自验证截图

【CheckList】

PR提交人对以下CheckList自检项进行全量自检,自检通过或不涉及,均修改 [ ] 为 [x]

likedislike
Pull Request已成功合入, 合并人@ascend-robot
(感谢 a_knight 的贡献)
Aa_knight
8月25日 创建了 pull request,commit 022b4c41
Aa_knight
8月25日 关联了issue:flexattention: An error occurs when seq_length is 180k
atomgit-bot
atomgit-bot
8月25日 评论:

变更摘要

该 PR 将 flex attention 的溢出修复移植到最新 master:针对反向传播模板中索引计算可能发生的整数溢出,引入 INDEX_DTYPE(tl.int32/tl.int64)统一索引类型;同时对 dK/dV tasklist reduce 配置增加 UB 预算过滤,并将稀疏块与 tile 块的匹配从「相等」放宽为「整除」以支持子块(sub-block)遍历,从而修复大 shape 或非对齐块场景下的溢出/错误问题。

主要改动

  • 索引溢出修复(INDEX_DTYPE): 在 flexattention_template.py 中将 tile_id、task_id 及各类偏移计算(原 .to(tl.int64))统一转换为 INDEX_DTYPE,并在 select_algorithm.py 的 NPUTritonTemplate 中新增 _write_index_dtype_define,允许 flex_attention 模板根据 can_use_32bit_indexing 选择 tl.int32 或 tl.int64,避免大 shape 下索引溢出。
  • dK/dV tasklist reduce 的 UB 预算过滤: 在 flex_attention.py 新增 _filter_dkdv_tasklist_reduce_configs,按 npu_config.ub_size 的固定比例计算 UB 预算并过滤 BLOCK_N1/head dim 过大的配置;若无配置满足约束,则将 tasklist_reduce_ub_safe 置为 false 并禁用该图的 tasklist 代码生成。
  • 子块(sub-block)支持: flex_attention_config_generator.py 将块匹配条件由 block_m == sparse_q_block_size 改为 sparse_q_block_size % block_m == 0,模板中对应将静态断言改为整除判断,并通过 SPARSE_Q_MULTIPLE/SPARSE_KV_MULTIPLE 及 kv_work_pos/kv_sub_block 循环遍历稀疏 KV 块内的子块。
  • dK/dV-only 任务数修正: flex_attention_backward_dkdv_only_source 中 NUM_TASKS 由 NUM_KV_BLOCKS * ZKV * HKV 修正为 NUM_KV_BLOCKS * ZQ * HKV,同时相关地址偏移(q_adj1、do_adj1、off_chz1 等)改用 INDEX_DTYPE。
likedislike
不准确?
atomgit-bot
atomgit-bot
8月25日 评论:

代码审查

✅ 未发现问题

likedislike
不准确?
此处折叠了160条消息 查看更多
ascend-robotascend-robot成员
8月26日 添加了label:lgtm
Xuan Peng
8月26日 评论:

/approve

likedislike
ascend-robotascend-robot成员
8月26日 添加了label:approved
ascend-robotascend-robot成员
8月26日 合入了pull request
ascend-robot
ascend-robot成员
8月26日 评论:
流水线 pytorch_gitcode_PR_multiVersion#14513 [ commitID:ed38019e ] 已完成
likedislike