跳到主要内容

Tensor 编程:MP31 TCE、WMMA、SQMMA 与 TME

本章介绍 MTCC 在 mp_31__MUSA_ARCH__ == 310)上提供的 Tensor builtin。 这些接口属于设备端编程模型;请使用本章列出的 MP31 shape、datatype 和同步 规则。不要把其他平台的 WMMA/SQMMA wrapper 或未声明的异步 wrapper 当作本章 API。

计算与搬运的分工

TCE 执行 D = A × B + C,TME 将 descriptor 描述的 global tile 搬到 shared memory,或把 tile 写回 global memory。MP31 operation 提供 layout、dtype、barrier 和 builtin 参数的组合;应用 kernel 应使用这些公开 operation,而不是重新拼接 raw 整数参数。

WMMA:warp MMA

MP31 warp-MMA operation 提供以下 7 个 builtin:

__musa_wmma_m16n8k4_mma__musa_wmma_m16n8k8_mma__musa_wmma_m16n8k16_mma__musa_wmma_m8n16k16_mma__musa_wmma_m16n16k16_mma__musa_wmma_m16n16k32_mma__musa_wmma_m16n16k64_mma

可编译的调用入口是 MTCC 的 MP31 operation。形状、A/B major mode、accumulator 和 TceABDtype 必须使用同一 operation 的定义;不要将其他架构的 fragment 签名直接套用到 MP31 builtin。

SQMMA:shape-specific MMA

MP31 SQMMA operation 使用以下 shape spelling:

m16n64m32n32m32n64m32n128m64n16m64n32m64n64m64n128m128n32m128n64m128n128

每个 family 的 builtin spelling、寄存器布局和 datatype selector 都由 MP31 operation 选择。不要从 WMMA 的参数顺序推导 SQMMA raw 调用。

TME descriptor 与 tile

descriptor、rank、维度、stride、global base address 和 OOB fill 的 host 编码应 遵循 MTCC descriptor API;device 侧 tile 操作使用 MP31 TME operation。MP31 使用 1D–5D tile load/store、3D–5D im2col load,以及 block load/store。所有异步 load 都必须配合 barrier transaction bytes 和 phase wait;预取-only builtin 只是提示, 不能替代 load 或完成同步。

Barrier 与流水

MP31 barrier ID 为 block-local;调用 __musa_async_arrive(barrier_id) 时, barrier_id 必须是 1–63,0 保留不可用。MTCC transaction barrier 和 MP31 pipeline 负责 arrival、phase、transaction accounting 和 stage reuse。生产者与 消费者必须保持一致控制流,且 transaction bytes 必须等于该 load 实际写入 shared memory 的字节数。

编译与验证

  1. 先运行普通 load/store 回退,验证边界、layout 和数值。
  2. 选择 MTCC MP31 MMA/TME operation,不要手写未声明的 builtin 参数。
  3. 使用 --offload-arch=mp_31 编译,并在目标设备上验证数值和边界。
  4. 对不支持的 shape、alignment 或架构保留普通搬运和计算回退。

完整 builtin 清单、原型、datatype selector、rank-specific load/store 和 prefetch 语义见 TME API Reference;流水线细节见 TME 与异步数据搬运