- [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]:不支持
接口功能:
实现MHC Post组件的前向计算,用于Transformer模型中多层残差连接的后处理阶段。该算子将残差矩阵变换与输出状态投影融合为单次计算,避免多次独立算子调用带来的额外开销。
计算公式:
MHC Post算子的核心计算公式为:
其中:
表示对输入进行残差矩阵的转置矩阵乘法。对于输出中的第行(对应第个head),计算过程为:
即将矩阵按转置方式与做矩阵乘法,为标量,对的第行做标量乘法后累加到第行输出。
表示输出状态与后处理权重的逐元素乘法与广播。对于第个head:
即为标量,对整行做标量乘法后加到第行输出。
综合完整计算过程为:
其中,表示参数
[object Object],表示参数[object Object],表示参数[object Object],表示参数[object Object],表示输出[object Object]。
[object Object]
[object Object]
该接口支持训练、推理场景下使用。
该接口支持单算子模式调用。
数据类型约束:
[object Object]和[object Object]的数据类型必须相同。- 输出
[object Object]的数据类型与[object Object]保持一致。
维度约束:
[object Object]的维度需与[object Object]维度格式匹配:4维时为(B, S, n, n),3维时为(T, n, n)。
Shape一致性约束:
- 4维(BSND)格式下:
[object Object]的(B, S)维度需与[object Object]的(B, S)维度一致,[object Object]的后两维为(n, n),其中n与[object Object]的第3维一致。[object Object]的(B, S)维度需与[object Object]的(B, S)维度一致,[object Object]的D维度需与[object Object]的D维度一致。[object Object]的(B, S)维度需与[object Object]的(B, S)维度一致,[object Object]的n维度需与[object Object]的n维度一致。
- 3维(TND)格式下:
[object Object]的T维度需与[object Object]的T维度一致,后两维为(n, n)。[object Object]的T维度需与[object Object]的T维度一致,D维度需与[object Object]的D维度一致。[object Object]的T维度需与[object Object]的T维度一致,n维度需与[object Object]的n维度一致。
- 4维(BSND)格式下:
所有输入Tensor的shape各维度值必须为正数(大于0)。
默认支持确定性计算。
单算子模式调用:
[object Object]