Pull Request已成功合入, 合并人@CANN-robot
(感谢 han-dongchen 的贡献)变更摘要
本次 PR 对 DSv4 系列 attention metadata 算子(共 6 个算子)的 host 端与 aicpu 端参数校验逻辑进行了统一增强。主要包括:新增 head_dim 必须为 128 的约束、将 cu_seqlens 的"元素非负"检查改为更严格的"首元素必须为 0"检查、为 seqused 系列参数新增上界校验(不超过对应 max_seqlen 或由 cu_seqlens 推导的序列长度)、对 ori_topk_length/cmp_topk_length 的维度一致性增加多维度校验,以及将 cmp_residual 错误信息中补充 cmpRatio 的实际值。部分函数签名因新增校验参数而发生变更。
主要改动
-
新增
head_dim == 128强制校验: 在CheckSingleParamLiV2和CheckSingleParamQliV2中增加了headDim参数,并校验其必须等于 128,不满足则返回参数无效错误。 -
cu_seqlens校验逻辑从「非负」改为「首元素为 0」: 在所有 aicpu kernel 的ParamsCheck中(涵盖cu_seqlens_q、cu_seqlens_k、cu_seqlens_ori_kv、cu_seqlens_cmp_kv),移除了逐个元素的>=0检查,改为检查首元素必须为 0。由于原有单调递增校验,首元素为 0 即可保证全部元素非负。 -
seqused参数新增上界校验: 在 aicpu kernel 端对seqused_q、seqused_k、seqused_ori_kv、seqused_cmp_kv增加上限校验——BSND 布局下不超过对应max_seqlen,TND 布局下不超过由cu_seqlens[i+1] - cu_seqlens[i]计算的序列长度。 -
ori_topk_length/cmp_topk_length维度一致性校验增强: 在sparse_flash_mla_metadata、sparse_flash_mla_grad_metadata、mixed_quant_sparse_flash_mla_metadata的 host 端CheckConsistency函数中,增加了对这两个张量各维度的校验——BSND 时校验第1维等于 batch、第2维等于max_seqlen_q、第3维等于num_heads_kv;TND 时校验第2维等于num_heads_kv。同时将校验触发条件从 socVersion(Ascend950)改为oriTopk/cmpTopk != 0且 mask mode 为DEFAULT_MASK。 -
BSND 布局下
max_seqlen必须 > 0 的校验: 在多个 host 端校验函数中新增规则:当layout_q或layout_kv为BSND时,对应max_seqlen_q/max_seqlen_k/max_seqlen_ori_kv/max_seqlen_cmp_kv必须大于 0,否则返回参数无效错误。


Thanks for your pull-request.
The full list of commands accepted by me can be found at here。
You can get sig-info at here
PR Approval Progress
✅ Congratulations! All modules have met the lgtm and approve requirements.
Module Approval Details
| module | lgtm status | approve status |
|---|---|---|
| attention | ✅ shasha_an, wangzhe123456789 (2/2) | ✅ wangzhe123456789 (1/1) |
💡 Tip:
- Committer can comment
/approveor/lgtm- Commenting
/approveimplies both code review (lgtm) and intent to merge (approve)
CLA Signature Pass
qq_32807861, thanks for your pull request. All authors of the commits have signed the CLA. 👍


描述
mqsmla & smla & smlag metadata:
增加l校验layout_q为BSND时,max_seqlen_q必须传入;
增加l校验has_ori_kv 且 layout_kv 为 BSND 时,max_seqlen_ori_kv 必须传入;
增加l校验has_cmp_kv 且 layout_kv 为 BSND 时,max_seqlen_cmp_kv 必须传入;
对ori_topk_length的校验添加先决条件:只在ori_topk不为0且ori_mask_mode为0时再校验;cmp_topk_length的校验同理;
增加l校验layout_q为BSND时,ori_topk_length三个轴的一致性校验;layout_q为TND时,ori_topk_length第二轴的一致性校验;cmp_topk_length同理;
增加l校验cu_seqlens_xxx首元素为0的校验;
增加l校验seqused_xxx元素不大于 max_seqlen_xxx (BSND) 或 cu_seqlens_xxx 序列长度的校验;
qli_v2 & li_v2 & sliklg metadata:
增加l校验layout_q为BSND时,max_seqlen_q必须传入;
增加l校验layout_k 为 BSND 时,max_seqlen_k 必须传入;
增加l校验cu_seqlens_xxx首元素为0的校验;
增加l校验seqused_xxx元素不大于 max_seqlen_xxx (BSND) 或 cu_seqlens_xxx 序列长度的校验;
关联的Issue
关联Issue #4087
测试
二级冒烟、编译触发
文档更新
类型标签