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 memory和shared 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 所需的卷积参数。它将 padding、stride 和 dilation 转换为设备端可使用的参数:
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 * 2,stride 的长度为 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_pos、output_dim 和 tme_conv_param。请根据搬运类型选择对应 overload,因为不同 overload 的尾参数具有不同语义。