当前位置: 首页 > news >正文

CANN稀疏FlashAttention反向算子

SparseFlashAttentionGrad

【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer

产品支持情况

产品是否支持
Ascend 950PR/Ascend 950DT
Atlas A3 训练系列产品/Atlas A3 推理系列产品
Atlas A2 训练系列产品/Atlas A2 推理系列产品
Atlas 200I/500 A2 推理产品×
Atlas 推理系列产品×
Atlas 训练系列产品×

功能说明

  • 算子功能:根据topkIndices对key和value选取大小为selectedBlockSize的数据重排,接着进行训练场景下计算注意力的反向输出。

  • 计算公式:根据传入的topkIndice对keyIn和value选取数量为selectedBlockCount个大小为selectedBlockSize的数据重排,公式如下:

    $$ selectedKey\text{ }=\text{ }Gather \left( key,topkIndices \left[ i \left] \left) ,\text{ }0\text{ } < =i < \text{ }selectBlockCount\right. \right. \right. \right. $$

    $$ selectedValue\text{ }=\text{ }Gather \left( value,topkIndices \left[ i \left] \left) ,\text{ }0\text{ } < =i < \text{ }selectBlockCount\right. \right. \right. \right. $$

阶段1:根据矩阵乘法导数规则,计算$dP$和$dV$:

$$ dP\mathop{{}}\nolimits_{{t,:}}=dO\mathop{{}}\nolimits_{{t,:}}\text{@}V\mathop{{}}\nolimits^{{T}} $$

$$ dV \left[ u \left] =P\mathop{{}}\nolimits_{{T}}^{{t,:}}\text{@}dO\mathop{{}}\nolimits_{{t,:}}\right. \right. $$

阶段2:计算$dS$:

$$ d\mathop{{S}}\nolimits_{{t,:}}= \left[ P\mathop{{}}\nolimits_{{t,:}}@ \left( dP\mathop{{}}\nolimits_{{t,:}}-FlashSoftmaxGrad \left( dO,O \left) \left) \right] \right. \right. \right. \right. $$

阶段3:计算$dQ$与$dK$:

$$ d\mathop{{Q}}\nolimits_{{t,:}}=d\mathop{{S}}\nolimits_{{t,:}}@K \left[ u \left] \mathop{{}}\nolimits_{{:t,:}}/\sqrt{{d\mathop{{}}\nolimits_{{k,:}}}}\right. \right. $$

$$ dK \left[ u \left] \mathop{{}}\nolimits_{{:t,:}}=dS\mathop{{}}\nolimits_{{t,:t}}\mathop{{}}\nolimits^{{T}}\text{@}Q/\sqrt{{d\mathop{{}}\nolimits_{{t,:}}}}\right. \right. $$

参数说明

参数名输入/输出/属性描述数据类型数据格式
query输入attention结构的输入Q。BFLOAT16、FLOAT16ND
key输入attention结构的输入K。BFLOAT16、FLOAT16ND
value输入attention结构的输入v。BFLOAT16、FLOAT16ND
sparseIndices输入稀疏场景下选择的权重较高的注意力索引。INT32ND
dOut输入注意力输出矩阵的梯度。BFLOAT16、FLOAT16ND
out输入注意力输出矩阵。BFLOAT16、FLOAT16ND
softmaxMax输入注意力正向计算的中间输出。FLOAT32ND
softmaxSum输入注意力正向计算的中间输出。FLOAT32ND
actualSeqLengthsQueryOptional输入每个Batch中,Query的有效token数。INT32ND
actualSeqLengthskvOptional输入每个Batch中,Key、value的有效token数。INT32ND
queryRopeOptional输入MLA rope部分:Query位置编码的输出。BFLOAT16、FLOAT16ND
keyRopeOptional输入MLA rope部分:Key位置编码的输出。BFLOAT16、FLOAT16ND
scaleValue属性缩放系数。FLOAT32-
sparseBlockSize属性选择的块的大小。INT64-
layout属性layout格式。STRING-
sparseMode属性sparse的模式。INT64-
preTokens属性Attention算子里, 对S矩阵的滑窗起始位置。INT64-
nextTokens属性Attention算子里, 对S矩阵的滑窗终止位置。INT64-
deterministic属性确定性计算。BOOL-
dQuery输出表示query的梯度。BFLOAT16、FLOAT16ND
dKey输出表示key的梯度。BFLOAT16、FLOAT16ND
dValue输出表示value的梯度。BFLOAT16、FLOAT16ND
dQueryRopeOptional输出表示queryRope的梯度。BFLOAT16、FLOAT16ND
dKeyRopeOptional输出表示keyRope的梯度。BFLOAT16、FLOAT16ND

约束说明

  • 参数query中的D和key、value的D值相等为512,参数query_rope中的Dr和key_rope的Dr值相等为64。
  • 参数query、key、value的数据类型必须保持一致。
  • 当前只支持value和key完全一致的场景。
  • 当前仅支持sparseMode=0或3(无mask或以右顶点为划分的下三角场景)
  • 仅支持BSND或TND layout;关于数据shape的约束如下:
    • B:取值范围1~256。
    • S1、S2:1~128K;S1、S2支持不等长。
    • N1支持1/2/4/8/16/32/64/128。
      • Ascend 950PR/Ascend 950DT :
        • 额外还支持48、24、12、6、3。
    • N2:仅支持1。
    • D:仅支持512。
    • Drope:仅支持64。
    • topk:1024、2048、3072、4096、5120、6144、7168、8192。
      • 不建议topk * sparseBlockSize超过100k,由于内部算法硬件限制可能会导致oom。
  • 确定性计算:
    • SparseFlashAttentionGrad默认非确定性实现,支持通过aclrtCtxSetSysParamOpt开启确定性。

调用示例

调用方式样例代码说明
aclnn接口test_aclnn_sparse_flash_attention_grad通过 aclnnSparseFlashAttentionGrad 接口方式调用算子

【免费下载链接】ops-transformer本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。项目地址: https://gitcode.com/cann/ops-transformer

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

http://www.jsqmd.com/news/849362/

相关文章:

  • Keil开发环境下的CANopen与DeviceNet协议实现指南
  • CANN/ops-blas Ssyr算子实现
  • NCE外汇:服务体验与平台稳定性的协同提升
  • svelte-preprocess 性能优化最佳实践:提升构建速度的10个技巧
  • CANN社区Sign算子优化设计
  • Spire性能优化技巧:如何高效使用Rational和SafeLong提升Scala数值计算效率
  • Element React终极指南:快速构建企业级React应用UI界面
  • HC-05蓝牙模块连接Arduino/STM32的3.3V/5V电平匹配全攻略,附电路图与代码
  • 如何在Windows11中安装Android应用?WSA工具使用教程
  • CANN AsNumpy排序函数API
  • ops-collections架构深度解析:如何实现NPU上的高性能哈希表
  • 别再被数学劝退!用PyTorch从零实现DDPM扩散模型(附完整代码)
  • 通过环境变量为hermesagent配置taotoken作为自定义模型服务提供方
  • CANN/asc-devkit 设置梯度输出类型
  • CANNBot torch-compile 快速入门
  • 2026河北钢制防火门多少钱一平米?甲乙丙级最新报价
  • CANN混元视频配置说明
  • 数据中心工频UPS哪家好?2026工频不间断电源/核磁用UPS电源生产厂家权威推荐 - 栗子测评
  • CTF中的音频隐写术实战:从‘兔耳’和‘调频收音机’两道Misc题,学会用Python脚本提取隐藏信息
  • HermesAgent工具连接Taotoken自定义模型提供方的完整流程
  • CANN Bench交叉熵损失算子评测
  • Matlab阶跃响应性能指标自动化计算:从原理到工程实践
  • 如何快速上手elec-ops-inspection:昇腾平台部署指南
  • Configor 自动重载功能深度解析:实现配置热更新的终极指南
  • CANN/hccl RDMA QP端口配置路径
  • 轨距调整片定制哪家好?2026年绝缘轨距块生产厂家优质供应商推荐指南:新建铁路配件领衔 - 栗子测评
  • 2026机房不间断电源生产厂家哪家好?深圳不间断电源生产厂家实力深度解析 - 栗子测评
  • cann/asc-devkit SetGradOutput接口
  • CANN ops-fft部署指南:生产环境中的配置、监控与故障排除
  • npc_gzip异常处理与调试手册:解决压缩器错误的10个实用技巧