已合并
update dsv4 metadata param check #9510
update dsv4 metadata param check #9510
已合并
han-dongchen创建于 8月3日
han-dongchen
han-dongchen成员
8月3日

描述

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

测试

二级冒烟、编译触发

文档更新

类型标签

likedislike
Pull Request已成功合入, 合并人@CANN-robot
(感谢 han-dongchen 的贡献)
han-dongchenhan-dongchen成员
8月3日 创建了 pull request,commit 3bb44c10
atomgit-bot
atomgit-bot
8月3日 评论:

变更摘要

本次 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 强制校验: 在 CheckSingleParamLiV2CheckSingleParamQliV2 中增加了 headDim 参数,并校验其必须等于 128,不满足则返回参数无效错误。

  • cu_seqlens 校验逻辑从「非负」改为「首元素为 0」: 在所有 aicpu kernel 的 ParamsCheck 中(涵盖 cu_seqlens_qcu_seqlens_kcu_seqlens_ori_kvcu_seqlens_cmp_kv),移除了逐个元素的 >=0 检查,改为检查首元素必须为 0。由于原有单调递增校验,首元素为 0 即可保证全部元素非负。

  • seqused 参数新增上界校验: 在 aicpu kernel 端对 seqused_qseqused_kseqused_ori_kvseqused_cmp_kv 增加上限校验——BSND 布局下不超过对应 max_seqlen,TND 布局下不超过由 cu_seqlens[i+1] - cu_seqlens[i] 计算的序列长度。

  • ori_topk_length/cmp_topk_length 维度一致性校验增强: 在 sparse_flash_mla_metadatasparse_flash_mla_grad_metadatamixed_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_qlayout_kvBSND 时,对应 max_seqlen_q/max_seqlen_k/max_seqlen_ori_kv/max_seqlen_cmp_kv 必须大于 0,否则返回参数无效错误。

likedislike
不准确?
CANN-robotCANN-robot成员
8月3日 添加了label:cann-cla/yes
CANN-robot
CANN-robot成员
8月3日 评论:

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 /approve or /lgtm
  • Commenting /approve implies 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. 👍

likedislike
han-dongchenhan-dongchen成员
8月3日 修改了pull request 的描述
此处折叠了215条消息 查看更多
CANN-robotCANN-robot成员
8月7日 添加了label:approved
shasha_an成员
8月7日 评论:

/lgtm

likedislike
CANN-robotCANN-robot成员
8月7日 添加了label:lgtm
CANN-robotCANN-robot成员
8月7日 关闭了关联的issue
CANN-robotCANN-robot成员
8月7日 合入了pull request