开发者
下载
[object Object][object Object]
  • [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]
  • 接口功能:

    实现MHC Post组件的前向计算,用于Transformer模型中多层残差连接的后处理阶段。该算子将残差矩阵变换与输出状态投影融合为单次计算,避免多次独立算子调用带来的额外开销。

  • 计算公式:

    • MHC Post算子的核心计算公式为:

      xl+1=(Hlres)Txl+hloutHtpostx_{l+1} = (H_{l}^{res})^{T} \cdot x_{l} + h_{l}^{out} \cdot H_{t}^{post}

      其中:

      • (Hlres)Txl(H_{l}^{res})^{T} \cdot x_{l} 表示对输入xlx_{l}进行残差矩阵的转置矩阵乘法。对于输出中的第ii行(对应第ii个head),计算过程为:

        xl+1[i]=j=0n1Hlres[j,i]xl[j]x_{l+1}[i] = \sum_{j=0}^{n-1} H_{l}^{res}[j, i] \cdot x_{l}[j]

        即将HlresH_{l}^{res}矩阵按转置方式与xlx_{l}做矩阵乘法,Hlres[j,i]H_{l}^{res}[j, i]为标量,对xlx_{l}的第jj行做标量乘法后累加到第ii行输出。

      • hloutHtposth_{l}^{out} \cdot H_{t}^{post} 表示输出状态hlouth_{l}^{out}与后处理权重HtpostH_{t}^{post}的逐元素乘法与广播。对于第ii个head:

        xl+1[i]+=Htpost[i]hloutx_{l+1}[i] += H_{t}^{post}[i] \cdot h_{l}^{out}

        Htpost[i]H_{t}^{post}[i]为标量,对hlouth_{l}^{out}整行做标量乘法后加到第ii行输出。

    • 综合完整计算过程为:

      xl+1[i,:]=Htpost[i]hlout[:]+j=0n1Hlres[j,i]xl[j,:]x_{l+1}[i, :] = H_{t}^{post}[i] \cdot h_{l}^{out}[:] + \sum_{j=0}^{n-1} H_{l}^{res}[j, i] \cdot x_{l}[j, :]

      其中,xlx_{l}表示参数[object Object]HlresH_{l}^{res}表示参数[object Object]hlouth_{l}^{out}表示参数[object Object]HtpostH_{t}^{post}表示参数[object Object]xl+1x_{l+1}表示输出[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]保持一致。
  • 维度约束:

    • [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维度一致。
  • 所有输入Tensor的shape各维度值必须为正数(大于0)。

[object Object]

默认支持确定性计算。

[object Object]
  • 单算子模式调用:

    [object Object]