- [object Object]Ascend 950PR/Ascend 950DT[object Object]:支持
- [object Object]Atlas A3 训练系列产品/Atlas A3 推理系列产品[object Object]:不支持
- [object Object]Atlas A2 训练系列产品/Atlas A2 推理系列产品[object Object]:不支持
- [object Object]Atlas 200I/500 A2 推理产品[object Object]:不支持
- [object Object]Atlas 推理系列产品[object Object]:不支持
- [object Object]Atlas 训练系列产品[object Object]:不支持
接口功能:
[object Object]是基于[object Object]的[object Object]扩展接口,用于调用[object Object]算子完成共享KV(Key和Value使用同一份输入)的非量化注意力计算,训练推理归一化。[object Object]是[object Object]的元数据生成接口,用于在主算子执行前生成metadata。metadata记录AICore/AIVCore的任务切分结果,主算子可选择传入该metadata以优化调度。典型调用流程如下:- 准备
[object Object]、[object Object]、[object Object]等输入。 - 调用
[object Object]生成[object Object]。 - 调用
[object Object],将上一步得到的[object Object]传入主算子。
- 准备
计算公式:
self-attention(自注意力)利用输入样本自身的关系构建了一种注意力模型。其原理是假设有一个长度为的输入样本序列,的每个元素都是一个维向量,可以将每个维向量看作一个token embedding,将这样一条序列经过3个权重矩阵变换得到3个维度为的矩阵。
self-attention的计算公式一般定义如下,其中为输入样本的重要属性元素,是输入样本经过空间变换得到,且可以统一到一个特征空间中。公式及算子名称中的"Attention"为"self-attention"的简写。
本算子中Score函数采用Softmax函数,self-attention计算公式为:
其中和的乘积代表输入的注意力,为避免该值变得过大,通常除以进行缩放,并对每行进行softmax归一化,与相乘后得到一个的矩阵。
增加sink之后计算逻辑如下所示,主要修改相关softmax_max和softmax_sum逻辑计算部分。
开启return_softmax_lse之后,返回值softmax_lse计算逻辑如下所示:
[object Object]
调用flash_attn接口之前,请先调用前置接口flash_attn_metadata,完成flash_attn负载均衡的计算。
说明
- attn_out:Tensor类型,公式中的输出,数据类型支持float16、bfloat16。数据格式支持ND。限制:该输出参数的D维度与value的D保持一致,其余维度需要与入参query的shape保持一致。
- softmax_lse:Tensor类型,ring attention算法对query乘key的结果,先取max得到softmax_max。query乘key的结果减去softmax_max,再取exp,最后取sum,得到softmax_sum,最后对softmax_sum取log,再加上softmax_max得到的结果。数据类型支持float32,return_softmax_lse为True时,一般情况下,输出shape为(B, Q_N, Q_S)的Tensor,当input_q为TND时,输出shape为(Q_N, Q_T)的Tensor;return_softmax_lse为False时,则输出shape为[1]的值为0的Tensor。
- 声明
- 参数cu_seqlens_q、cu_seqlens_kv、seqused_q、seqused_kv、block_table及attn_mask属于tensor。由于算子在Tiling阶段无法获取tensor的具体数值,tiling侧不对值进行校验,正确性需要用户自行保证。若上述参数传入非法值,会触发未定义行为(精度问题、非法内存访问导致的程序崩溃等)。
- flash_attn_metadata和flash_attn的入参在调用时应该保持一致。由于算子分为两个接口分段调用,算子无法自行校验,正确性需要由客户自行保证。若接口传入参数不一致,会发生未定义行为(精度问题、非法内存访问导致的程序崩溃等)。
资料约束中,常见字段释义如下:
入参为空的场景处理:
- 空Tensor指必选输入和输出的shape size为0,即有任意轴为0。
- 触发空tensor的用例将全部拦截报错。
q、k、v、attn_out校验
layout匹配关系表:
[object Object]metadata校验
[object Object]mask_mode参数解释
[object Object][object Object]当block_table不为空时,开启Paged Attention
[object Object]flash_attn_metadata + flash_attn 联合调用示例(BSND)
[object Object]