跳到主要内容

TME/TCE API 参考

概述

本页提供适用于 S5000(编译目标 mp_31)的 TCE(Tensor Core Engine,张量计算引擎)和 TME(Tensor Memory Engine,张量内存引擎)公开 C/C++ API 参考,涵盖 descriptor、异步数据搬运、barrier 和矩阵计算接口,并说明接口原型、参数含义、调用顺序和使用约束。具体支持情况以当前 SDK 的公开头文件和示例为准;完整的 Tensor 编程流程请参见 Tensor 编程

MUSA Tensor API 由两个相互配合的部分组成:

  • 主机端代码(Host Code):创建 tensor descriptor,并编码 im2col 或 direct CONV 所需的参数;
  • 设备端代码(Device Code):选择 tile,发起 TME 搬运,使用 barrier 等待数据就绪,再调用 WMMA、SQMMA 或 CONV 完成计算。

descriptor、tile 参数、barrier 和 TCE fragment 分别描述不同层级的对象。

通用规则

  • global memoryshared memory 分别表示设备全局内存和线程块共享内存。
  • ldm 表示 leading dimension;WMMA、SQMMA 和 CONV 的 load/store API 以元素数为单位。
  • descriptor stride 和 make_sqmma_desc 的 leading stride 以字节数为单位。
  • fragment 是由参与线程协同使用的分布式寄存器对象。
  • descriptor 创建成功只表示主机端编码有效;设备端仍需使用受支持的 tile、layout、shape 和数据类型组合。
  • TME 搬运和 TCE 计算是独立阶段,需要分别按 barrier 或 group 规则同步。
  • WMMA、SQMMA 和 CONV 等协作式 API 必须由参与线程以一致的控制流调用。

主机端 descriptor API

muTensorDescriptorEncode

在 Host Code 中使用 muTensorDescriptorEncode 描述全局 tensor,并生成供 TME 使用的 MUtensorDescriptor。它编码设备地址、元素类型、维度、字节 stride、interleave 和越界填充值:

MUresult muTensorDescriptorEncode(
MUtensorDescriptor *desc, MUtensorDescriptorDataType type,
uint32_t rank, void *device_ptr, const uint64_t *dims,
const uint64_t *strides, MUtensorDescriptorInterleave interleave,
uint64_t oob_fill);

参数和结果约束如下:

  • dims 的最连续维在前,长度至少为 rank
  • strides 的单位是字节,最连续维不需要 stride,长度至少为 rank - 1
  • 普通 TME tile、im2col 和 swizzle 使用 MU_TENSOR_DESCRIPTOR_INTERLEAVE_NONE
  • 调用成功后,descriptor 可以按值传入 kernel,也可以复制 64 字节到设备后传递设备端指针;
  • 返回值为 MUSA_SUCCESS 后,才继续使用 descriptor。

muTensorIm2colConvParamEncode

在 Host Code 中使用 muTensorIm2colConvParamEncode 编码 im2col 所需的卷积参数。它将 paddingstridedilation 转换为设备端可使用的参数:

MUresult muTensorIm2colConvParamEncode(
MUconvParamer *p, uint32_t rank, const uint32_t *padding,
const uint32_t *stride, const uint32_t *dilation);

三个数组的长度必须与 rank 一致。编码成功后,使用 __musa::tme_conv_param(p.params[0], p.params[1], p.params[2]) 将结果转换为设备端参数,并将其作为 im2col memcpy_async 的附加参数。

muTensorDirectConvEncode

在 Host Code 中使用 muTensorDirectConvEncode 编码 direct CONV 的 input/weight 维度、stride、padding、interleave、swizzle 和 basic shape:

MUresult muTensorDirectConvEncode(
MUdirectConvParamer *param, uint32_t conv_rank,
const uint32_t *input_dim, const uint32_t *input_stride,
const uint32_t *weight_dim, const uint32_t *weight_stride,
const uint32_t *padding, const uint32_t *stride,
MUdirectConvInterleave input_interleave,
MUdirectConvInterleave weight_interleave,
bool input_swizzle, bool weight_swizzle,
uint32_t fill, MUdirectConvBasicShape shape);

padding 的长度为 conv_rank * 2stride 的长度为 conv_rank。生成的 MUdirectConvParamer 必须与 input/weight 的元素类型、内存布局和 CONV shape 匹配;返回值成功后,再将编码结果传给 Device Code。


TME 同步和搬运 API

__musa::async_barrier

使用 __musa::async_barrier 管理异步 TME 搬运的到达、phase 和完成等待。下面的调用形态展示初始化 arrival、记录 phase 和等待数据就绪的基本顺序:

__musa::async_barrier bar(1);
bar.init_arrival(n);
unsigned phase = bar.arrive();
bar.wait(phase);

init_arrival 注册参与异步操作的到达次数;arrive 发布当前 phase;wait 保证异步 copy 已完成。memcpy_async 会根据 bytes 自动增加 transaction count;负责读取并处理本次 tile 的计算线程(或线程组)完成后,调用 dec_trans(bytes) 释放对应的传输字节数。arrival count 统计到达次数,transaction count 统计待完成的传输字节数,两者分别维护。读取共享内存前必须调用 wait,且初始化、发起 copy 和 arrive 的线程分工必须保持一致。

__musa::memcpy_async

使用 __musa::memcpy_async 根据 tensor descriptor 将全局 tensor 的 tile 异步搬运到共享内存,并通过 barrier 通知搬运完成:

__musa::memcpy_async(bar, smem_dst, desc_ptr,
block_dim.get_param(), block_pos.get_param(), bytes,
extra0, extra1, extra2);

bytes 是字节数;block_dim 是 tile 长度,block_pos 是全局起始坐标。普通 tile 形式使用布局控制参数;prefetch 形式增加预取字节数;im2col 形式增加 weight_posoutput_dimtme_conv_param。请根据搬运类型选择对应 overload,因为不同 overload 的尾参数具有不同语义。

__musa::memcpy

使用 __musa::memcpy 将共享内存中的 tile 按 descriptor 写回全局 tensor:

__musa::memcpy(smem_src, desc_ptr,
block_dim.get_param(), block_pos.get_param());

该接口执行 TME tile store。它不负责等待异步 load;异步 load 仍需通过 barrier wait 确认完成。

__musa::memcpy_idf_l2

TME store 发起后,可以使用 __musa::memcpy_idf_l2 处理 L2 数据流水和可见性:

__musa::memcpy_idf_l2();

这是无参数的 L2 数据流水和可见性控制点,应放在 TME store 后使用。它不承担 async barrier 的完成等待。

tme_block_dim_v2/v3/v4/v5tme_block_pos_v2/v3/v4/v5

发起 TME 搬运前,使用 tme_block_dim_vN 描述各维搬运长度,使用 tme_block_pos_vN 描述各维起点,并将结果传给对应的搬运 overload:

__musa::tme_block_dim_v3 dim(x,y,z);
__musa::tme_block_pos_v3 pos(px,py,pz);
auto d = dim.get_param();
auto p = pos.get_param();

dim 表示各维搬运长度,pos 表示各维起点;版本号必须匹配参数维数和 descriptor rank。多维形式提供 v2~v5,rank-1 使用标量参数形式。

__musa::tme_conv_param

使用 __musa::tme_conv_param 将 Host Code 编码的 im2col 参数包装为 Device Code 可传递的参数:

__musa::tme_conv_param p(words[0], words[1], words[2]);
auto raw = p.get_param();

这是设备侧 im2col 参数包装器;三个 word 必须来自 muTensorIm2colConvParamEncode


WMMA API

wmma::fragment

使用 wmma::fragment 声明 WMMA 的 A/B 操作数 fragment 或 accumulator fragment:

wmma::fragment<Use,M,N,K,Element,Layout> f;

Usematrix_amatrix_baccumulator;A/B 需要 row/col layout,accumulator 不带 layout。fragment 是 warp 协作对象,参与同一操作的线程必须一致声明。

wmma::load_matrix_sync

使用 wmma::load_matrix_sync 将 global memory 中的 A/B 矩阵装载到对应 fragment:

wmma::load_matrix_sync(frag, ptr, ldm);

协同装载 A/B;ldm 是元素数,不是字节。A row/col 通常为 K/M,B row/col 通常为 N/K。装载 shape/type 必须在主手册 WMMA 支持表中。

wmma::fill_fragment

在第一次 WMMA 矩阵乘加前,使用 wmma::fill_fragment 初始化 accumulator:

wmma::fill_fragment(c, 0.0f);

将 accumulator 的每个 lane 片段初始化为同一值;第一次 mma_sync 前必须调用,除非明确继续已有累加器。

wmma::mma_sync

使用 wmma::mma_sync 执行一次 warp 协作的矩阵乘加:

wmma::mma_sync(c_out, a, b, c_in, false);

计算 c_out = A * B + c_in。A/B fragment 的 shape、类型和 layout 必须与 accumulator 完全匹配;该调用是 warp-synchronous。

wmma::store_matrix_sync

使用 wmma::store_matrix_sync 将 accumulator 中的结果协同写回 global memory:

wmma::store_matrix_sync(dst, c, ldm, wmma::mem_row_major);

协同写回 accumulator;输出 ldm 仍按元素计数。mem_row_majormem_col_major 决定输出步长和布局。16x8x4 TF32 支持 row-major 输出;其他输出布局按开发者指南的支持组合选择。


SQMMA API

sqmma::sqmmadescsqmma::make_sqmma_desc

TME 将 tile 搬入 shared memory 并完成同步后,使用 sqmma::make_sqmma_desc 为 SQMMA 创建 A/B 操作数 descriptor:

sqmma::sqmmadesc<sqmma::desc_a,T> a;
sqmma::make_sqmma_desc(a, smem, stride_bytes, sqmma::sg_16_byte);

descriptor 描述共享内存 tile,不描述全局 tensor。stride_bytes 是字节;sg_* 是共享内存 swizzle granularity;A/B 的 descriptor 角色应与实际操作数保持一致。构造 descriptor 前必须完成 TME load 和线程同步。

sqmma::fragmentsqmma::fill_fragment

使用 sqmma::fragment 声明 accumulator,并在计算前使用 sqmma::fill_fragment 初始化它:

sqmma::fragment<sqmma::accumulator,M,N,K,AccT> c;
sqmma::fill_fragment(c, 0);

fragment 是分布式 accumulator;fill_fragment 初始化每个 lane 的寄存器片段,不代表初始化共享内存。

sqmma::fragment<sqmma::accumulator,M,N,K,T>T 是累加器类型:FP16/BF16/TF32 输入使用 float,INT8/UINT8 输入使用 intM/N/K 必须来自支持矩阵。

sqmma::mma_sync

使用 sqmma::mma_sync 读取 A/B descriptor,结合 accumulator 执行 SQMMA 矩阵乘加:

sqmma::mma_sync(out, a_desc, b_desc, c,
sqmma::mem_row_major, sqmma::mem_col_major,
sqmma::positive, sqmma::positive, sqmma::init_none);

a_layout/b_layout 解释 descriptor 背后的 A/B 逻辑布局;positive/negative 表示输入 scale;init_zero 从零开始,init_none 使用传入 C。支持 A row/col、B row/col 与输出 row/col 的组合。

sqmma::store_matrix_sync

使用 sqmma::store_matrix_sync 将 SQMMA accumulator 协同写回 global memory:

sqmma::store_matrix_sync(dst, out, ldm, sqmma::mem_row_major);

把 accumulator 写回全局内存;ldm 是元素数,布局参数决定目标内存解释。


CONV API

conv::convdescconv::make_conv_desc

TME 将 input/weight tile 搬入 shared memory 并完成同步后,使用 conv::make_conv_desc 创建 CONV 运算数 descriptor:

conv::convdesc<conv::desc_a,T> a;
conv::make_conv_desc(a, smem_a);

该 descriptor 包装 TME 搬入 shared memory 的 input/weight tile。空间、通道和 interleave 的解释来自 Host Code 生成的 direct CONV descriptor。

conv::fragmentconv::fill_fragmentconv::conv_sync

使用 conv::fragmentconv::fill_fragment 准备 accumulator,再使用 conv::conv_sync 执行 direct convolution:

conv::fragment<conv::accumulator,M,N,K,AccT> c;
conv::fill_fragment(c, 0);
conv::conv_sync(out, a, b, c, p.get_conv_param(),
conv::positive, conv::positive, conv::init_none);

p 应由 __musa::conv_param 构造;conv_sync 完成 direct convolution,两个 scale 参数和 init_none/init_zero 的语义与 SQMMA 相同。

conv::fragment<conv::accumulator,M,N,K,T> 是分布式输出累加器;FP16/BF16/TF32 输入使用 float,INT8/UINT8 输入使用 int

conv::store_matrix_sync

使用 conv::store_matrix_sync 将 CONV accumulator 协同写回 global memory:

conv::store_matrix_sync(dst, out, ldm, conv::mem_col_major);

协同写回 CONV accumulator。输出可以选择 row-major 或 col-major;shape 和类型组合必须来自开发者指南的 CONV 支持组合。


Prefetch API

prefetch

如果需要提前提示 LSU 读取后续数据,可以使用 prefetch

prefetch(&x[1], 64);

这是 LSU 预取,支持 int/float 的 64、128、256 字节形式。TME prefetch 没有独立函数,而是通过 memcpy_async 重载的尾参数传递。


相关文档

  • Tensor 编程:介绍 Tensor 编程模型、路径选择、开发流程和完整示例。
  • 协作组:介绍通用线程组同步、归约、扫描和异步搬运;需要 Tensor descriptor 或 TME barrier 时,请回到本页和 Tensor 编程页面。