跳到主要内容

Tensor 编程指南

概述

本指南以 S5000(编译目标 mp_31)为例,介绍 Tensor 编程路径,帮助你选择合适的 API,并完成数据搬运、同步和计算。

本文重点介绍:

  • 如何选择 WMMA、SQMMA、CONV 和 TME;
  • 如何使用 TME 在全局内存和共享内存之间搬运数据;
  • 如何使用 barrier 同步异步操作;
  • 如何使用 TCE API 执行矩阵乘加和卷积;
  • 如何将计算结果写回内存。

完整的 API 签名、参数说明以及支持的类型和 shape 组合,请参见 TCE/TME API 参考

适用场景

如果你准备在 Device Code 中自定义矩阵乘加或 direct convolution,可以先根据以下需求判断本指南是否适合:

  • 需要显式控制 tile 的 shape、数据类型或内存布局;
  • 需要控制 global memory 与 shared memory 之间的数据搬运和同步;
  • 需要自定义计算流程,或将矩阵、卷积计算与其他操作融合。

对于标准 GEMM、卷积或 attention,请先评估 MUSA 性能库是否已经提供所需能力。当现有算子无法满足数据布局、数据流或融合需求时,再使用本指南介绍的 Tensor Device Code 编程接口。

第一次编写 Tensor kernel 时,建议先实现并验证单个 tile 的“描述 → 搬运 → 同步 → 计算 → 写回”流程。结果正确后,再逐步加入 pipeline、double buffering、prefetch、swizzle 以及数据搬运与计算重叠。

编程模型

张量计算引擎(Tensor Core Engine, TCE)负责执行矩阵计算;张量内存引擎(Tensor Memory Engine, TME)负责在 global memory 和 shared memory 之间搬运 tile。 这两个引擎配合完成 Tensor kernel 的基本流程:TME 准备 tile,barrier 确认数据就绪,TCE 执行计算,最后将结果写回。

基本执行流程

在 S5000 上,一个 Tensor 程序通常按以下流程组织:

  1. 主机端描述数据:用描述符(descriptor)描述全局 tensor;im2col 和 direct CONV 还需要编码卷积参数。
  2. TME 搬运:设备端把全局内存的 tile 搬到共享内存,并用屏障等待搬运完成。
  3. TCE 计算:WMMA 直接从全局内存装载 fragment;SQMMA 和 CONV 从共享内存描述符读取 tile。
  4. 输出写回:使用对应的 store_matrix_sync,或使用 TME 写回全局内存。

下图展示三种计算路径在数据准备和计算阶段的差异:WMMA 从 global memory 装载 fragment;SQMMA 和 CONV 先通过 TME 将 tile 搬到 shared memory,再执行计算。

其中,shared memory 用于保存 SQMMA 和 CONV 的操作数,TCE 计算前必须完成数据填充和线程同步;fragment 是分布式寄存器对象,只能由参与同一 warp 或 block 的线程协同操作。

WMMA、SQMMA 和 CONV 共享上述基本流程,但在操作数来源、descriptor 和计算接口上有所不同。

计算路径

本指南介绍以下三种计算路径:

  1. WMMA(Warp-level Matrix Multiply-Accumulate,warp 级矩阵乘加):由 warp 内线程协同执行矩阵乘加。它从 global memory 装载矩阵数据到 fragment,完成计算后写回结果,适合需要控制 fragment、数据布局和较小矩阵 tile 的自定义 kernel。
  2. SQMMA(Shared Queue Matrix Multiply-Accumulate,面向共享内存 tile 的矩阵乘加):先通过 TME 将输入 tile 搬到 shared memory,使用 barrier 确认数据就绪,再从 tile 创建 A/B descriptor 并执行矩阵乘加,适合需要显式组织 tile 搬运、共享内存布局和计算流程的自定义 kernel。
  3. CONV(Convolution,卷积):执行 direct convolution。它先通过 TME 将 input 和 weight tile 搬到 shared memory,再结合 Host Code 编码的卷积参数创建 descriptor 并执行卷积,适合需要控制卷积 tile、数据布局或融合逻辑的自定义 kernel。

本指南只介绍公开 API。请通过对应 API 传递 layout、dtype、barrier 和 transaction 参数,并以接口定义的参数形式为准。具体接口签名和参数约束请参见 TCE/TME API 参考

基本概念

下表说明 Tensor kernel 中几个基本概念及其作用:

概念在 Tensor kernel 中的作用
tensor位于全局内存中的多维数据集合。
tile一次搬运或计算所处理的数据区域,由 shape 和 position 确定。
descriptor描述全局 tensor 的地址、类型、维度、stride 和相关布局信息。
shared memory保存 SQMMA 或 CONV 操作数的线程块共享存储空间。
fragment由参与计算的线程协同使用的分布式寄存器对象。
barrier通知线程异步 TME 搬运何时完成。

头文件与命名空间

下表列出各项能力对应的公开头文件、命名空间、操作数位置和适用场景:

能力头文件命名空间操作数位置适用场景
WMMA<mma.h>mtmusa::wmma全局内存 -> warp fragment小矩阵 tile
SQMMA<sqmma.h>mtmusa::sqmmaTME -> 共享内存 descriptor大矩阵 tile
CONV<conv.h>mtmusa::convTME -> 共享内存 descriptordirect convolution
TME<musa.h>__musatensor descriptor -> 共享/全局内存多维异步搬运、im2col 和预取

请直接 include 表中列出的公开头文件;这些头文件已经包含所需的底层声明,无需另行 include crt/*.h


运行原理

异步执行流程

  1. 初始化 arrival barrier。
  2. 为当前 TME load 记录准确的 transaction bytes。
  3. 发起 tile、block 或 im2col 搬运。
  4. 调用 barrier 的 arrive() 并保存返回的 phase。
  5. 在读取 shared memory 之前等待该 phase。
  6. 对载入的 tile 进行计算。
  7. 计算线程完成对 tile 的读取和处理后,减少对应的 transaction count。
  8. 计算线程完成对 stage 的释放后,stage 才可以复用。

Barrier 同步流程

发起 TME 异步搬运后,目标 shared memory 不能立即读取。参与线程需要按约定完成 arrival,保存返回的 phase,并在读取数据前等待该 phase。arrival count 统计需要到达的次数,transaction count 统计待完成的传输字节数,两者分别维护。

Tensor 数据搬运流程

当你需要把 global memory 中的 tensor tile 搬到 shared memory,并交给 WMMA、SQMMA 或 CONV 计算时,请按以下流程执行:

  1. 在主机端描述全局内存中的 tensor,并准备 descriptor。
  2. 在设备端选择 tile 的 shape 和 position。
  3. 使用 TME 将 tile 从全局内存搬到共享内存。
  4. 使用 barrier 等待异步搬运完成。
  5. 由 WMMA、SQMMA 或 CONV 读取并计算 shared memory 中的 tile,或使用 TME 将结果写回全局内存。

TME load 是异步操作,发起搬运不等于目标数据已经可读。读取目标共享内存前必须等待对应 phase;TME store 和异步 load 的 barrier wait 也属于不同的操作路径。

完成这个流程后,可以继续查看后面的 descriptor、tile 参数、TME API、Tensor 计算和高级搬运章节。

调用顺序

公开 wrapper 的基本调用顺序如下:

下图展示一次 TME 异步搬运的调用时序。TME 写入 shared memory 与 Device Code 等待数据就绪是两个并行参与的过程;只有 wait(phase) 返回后,才可以读取目标 tile。

  1. 初始化 __musa::async_barrier
  2. 使用 __musa::memcpy_async 发起 TME 搬运。
  3. 调用 arrive() 保存 phase。
  4. 调用 wait(phase),确认数据就绪后读取 shared memory。
  5. 完成 tile 的读取和处理后,调用 dec_trans(bytes) 释放对应的 transaction count。

具体接口签名和参数组合请参见 TME 同步和搬运 API


快速开始

按照以下步骤完成一个单 tile kernel,并验证计算结果。

准备工作

开始前,请准备:

  • 一台可用的 S5000 设备和对应的 MUSA SDK;
  • 输入 tensor 的设备地址、数据类型、维度和存储布局;
  • 要处理的 shape,以及 tile 的大小和位置;
  • 一个用于对照结果的 CPU 或普通 Device Code 实现;
  • 公开头文件或 SDK 示例中已确认支持的 API、shape 和类型组合。

选择计算路径

下表按计算路径说明其适用需求和数据准备方式:

计算路径适用需求数据准备方式
WMMA直接从 global memory 装载矩阵 fragmentWMMA load API
SQMMA对已经搬入 shared memory 的 tile 执行矩阵计算TME load → shared memory descriptor
CONV对搬入 shared memory 的 input/weight tile 执行卷积TME load → convolution descriptor

如果你要实现标准 GEMM、卷积或 attention,请先评估性能库。需要自定义数据流、布局或融合逻辑时,再使用这里的 Device Code 路径。

基础流程

  1. 在主机端代码(Host Code)中创建并检查 Tensor descriptor;im2col 或 direct CONV 还需要准备对应的卷积参数。
  2. 在设备端代码(Device Code)中根据 tensor rank 选择 tile shape 和 tile position。
  3. 为目标 shared memory 准备空间,并初始化 barrier。
  4. 记录本次搬运的 transaction bytes,发起 TME tile、block 或 im2col 搬运。
  5. 调用 arrive() 保存 phase,并在读取目标 shared memory 前调用 wait()
  6. 使用已确认支持的 WMMA、SQMMA 或 CONV 组合完成计算。
  7. 使用对应的 store API 或 TME 将结果写回 global memory。

每完成一步就检查其输出:先验证 descriptor 和 tile,再确认数据已经可读,最后验证计算和写回结果。基础流程正确后,再加入多个 stage、pipeline、prefetch 或 swizzle。

验证结果

运行 kernel 后,请确认:

  • 输出结果与 CPU 或普通 Device Code 实现一致;
  • barrier 的 arrival count 与实际参与者数量一致;
  • arrival count 与 transaction count 分别按到达次数和传输字节数维护;
  • ldm、descriptor stride 和 transaction bytes 使用了各自正确的单位;
  • 不支持的 shape、alignment 或边界情况保留普通路径回退。

完成基础流程后,请按以下任务继续:

  1. 创建 tensor descriptor 时,查看描述 Tensor
  2. 搬运和同步 tile 时,查看TME 数据搬运
  3. 选择矩阵乘加或卷积计算路径时,查看 WMMASQMMACONV 章节。

描述 Tensor

本节介绍如何在 Host Code 中准备 Tensor 和卷积参数,为后续的 TME 数据搬运以及 SQMMA、CONV 计算提供输入。请先完成描述并检查参数,再在 Device Code 中选择 tile 和搬运接口。

muTensorDescriptorEncode

muTensorDescriptorEncode 在 Host Code 中将设备地址、元素类型、tensor 的 rank、各维长度、字节 stride、interleave 和越界填充值编码为 MUtensorDescriptor。生成的 descriptor 用于描述 TME 要访问的全局 tensor。下面的示例描述一个 8 × 8int tensor:

MUtensorDescriptor desc;
const uint64_t dims[2] = {8, 8};
const uint64_t strides[1] = {8 * sizeof(int)};

MUresult rc = muTensorDescriptorEncode(
&desc, MU_TENSOR_DESCRIPTOR_DATA_TYPE_INT32,
/*rank=*/2, device_ptr, dims, strides,
MU_TENSOR_DESCRIPTOR_INTERLEAVE_NONE,
/*oob_fill=*/0);
  • 本例描述一个 8 × 8int tensor。dims 按最连续维在前排列,数组长度至少为 rank
  • strides 的单位为字节;最连续维不需要单独提供 stride,因此数组长度至少为 rank - 1。在本例中,strides[0] 表示跨过一个非连续维所需的字节数,即 8 * sizeof(int)
  • 普通 TME tile、im2col 和 swizzle 使用 MU_TENSOR_DESCRIPTOR_INTERLEAVE_NONE
  • 调用后应检查 rc 是否为 MUSA_SUCCESS。只有 descriptor 创建成功后,才应将其传给 kernel。
  • 创建成功后,descriptor 可以按值传入 kernel;也可以将其 64 字节内容复制到设备内存,再将设备端地址传给 kernel。具体方式应与对应的 TME API 参数形式保持一致。

完成 descriptor 创建后,请根据实际计算路径继续准备 im2col 或 direct CONV 所需的卷积参数。

im2col 卷积参数

如果使用 im2col 路径,请使用 muTensorIm2colConvParamEncode 在 Host Code 中编码 paddingstridedilation。设备端再用 __musa::tme_conv_param 包装编码结果,作为 im2col memcpy_async 的附加参数。下面的示例展示一维卷积参数的编码和设备端包装方式:

MUconvParamer host_param;
uint32_t padding[1] = {0};
uint32_t stride[1] = {1};
uint32_t dilation[1] = {1};
muTensorIm2colConvParamEncode(
&host_param, 1, padding, stride, dilation);

__musa::tme_conv_param device_param(
host_param.params[0], host_param.params[1], host_param.params[2]);

rank 决定这三个数组的维数;数组长度和卷积维度必须一致。调用后应检查返回值,并确认生成的参数与 input tensor 的实际布局和 im2col 搬运路径匹配。

direct CONV 参数

如果使用 direct CONV 路径,请使用 muTensorDirectConvEncode 编码 input 和 weight 的维度、stride、padding、interleave、swizzle 以及 basic shape。

padding 的长度为 conv_rank * 2stride 的长度为 conv_rank。调用后应检查返回值,并确认生成的 MUdirectConvParamer 与 input/weight 的元素类型、实际内存布局和 CONV shape 匹配。调整计算 shape 或内存布局后,请重新编码对应参数。

完成这些参数准备后,请继续阅读 TME 数据搬运,了解如何选择 tile、发起异步搬运并等待数据就绪;如果使用 direct CONV,还需要继续阅读 CONV 了解 Device Code 中的计算流程。


TME 数据搬运

本节介绍如何使用 TME 选择数据块、发起异步搬运、等待数据就绪,并将结果写回全局内存。开始前,请先完成 描述 Tensor 中的 descriptor 和相关卷积参数准备。

初始化 Barrier 并等待搬运完成

下面的代码展示一次异步 TME 搬运所需的 barrier 生命周期:初始化 barrier,发起搬运,记录 phase,等待数据就绪后再读取 shared memory。该代码用于说明调用顺序,不是完整的线程分工示例。

__musa::async_barrier bar(1);
bar.init_arrival(2);
__syncthreads();

// 一个或多个 warp 发起 memcpy_async
__musa::memcpy_async(bar, smem, desc_ptr,
dim.get_param(), pos.get_param(),
transfer_bytes, 0, 3, 1);

unsigned phase = bar.arrive();
bar.wait(phase);
// 现在才读取 smem

init_arrival(2) 要求两个参与者分别完成 arrival;实际 kernel 应让两个参与者分别调用 arrive(),或者将 arrival count 设置为实际参与者数量。init_arrival 注册 arrival count,arrive 发布当前 phase,wait 等待 TME 写入完成。barrier 初始化、copy 发起和 arrive 的线程分工必须保持一致;读取目标 shared memory 前必须完成 wait

tme_block_dimtme_block_pos

发起 TME 搬运前,需要分别指定 tile 的尺寸和在 tensor 中的起始位置。下面的代码将这两个参数编码为 memcpy_async 所需的设备端参数:

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

dim 表示每一维的搬运长度,pos 表示 tile 在 tensor 中的起始坐标。v2v3v4v5 分别对应不同的参数维数;选择的版本必须与 tensor rank 以及 memcpy_async overload 匹配。

__musa::memcpy_async

准备好 barrier、descriptor、tile 尺寸和起始位置后,使用 __musa::memcpy_async 将数据从全局内存异步搬运到 shared memory:

__musa::memcpy_async(bar, shared_dst, descriptor_ptr,
packed_dim, packed_pos,
bytes, extra0, extra1, extra2);

这是从全局内存到 shared memory 的异步 TME load。bytes 使用字节作为单位,必须等于目标 tile 的元素数量乘以单个元素的字节数。

  • 普通 tile:尾部参数使用对应 tile case 的 layout 参数,常见形式为 0, 3, 1
  • TME prefetch:尾部追加预取字节数。
  • im2col:尾部追加 weight position、output dimension 和 tme_conv_param

这三种参数形式对应不同的硬件路径,请根据搬运类型选择匹配的参数形式。

选择 overload 时,请先确认以下三项:

  • 你搬运的是普通 tile、block 还是 im2col 数据;
  • tile 的 rank、shape 和 position 是否与 descriptor 一致;
  • bytes 是否等于目标 tile 的元素数乘以元素大小。

如果有任何一项无法确认,请先回到“描述 Tensor”检查 descriptor 和 rank;确认数据路径后,再选择与当前场景匹配的 overload 和尾部参数。

__musa::memcpy__musa::memcpy_idf_l2

当计算线程完成对 shared memory 中 tile 的读取和处理后,可以使用以下接口将 tile 写回全局内存:

__musa::memcpy(shared_src, descriptor_ptr,
packed_dim, packed_pos);
__musa::memcpy_idf_l2();

__musa::memcpy 按 descriptor 将 shared memory 中的 tile 写回全局内存。__musa::memcpy_idf_l2 是 TME store 后的 L2 流水和可见性控制点;异步 load 仍需单独完成 barrier wait。

TME tile、im2col、swizzle

根据要搬运的数据和目标布局选择对应的 TME 路径:

  • tile:支持 rank 1~5,使用对应的 tme_block_dim_vNtme_block_pos_vN
  • im2col:使用 tme_conv_param,设备端 copy 还需要 weight position 和 output dimension。
  • swizzle:rank-1 tile load/store 使用 (SG, SS, SL) 配置;SG_NONE/16B/32B/64BSS_32B/64B/128B/256BSL_128B/256B 的组合必须与目标 tile 的连续字节数匹配。

选择 TCE API

TME 负责准备或写回数据,TCE 负责执行计算。WMMA、SQMMA 和 CONV 对操作数位置、descriptor、layout、累加器和线程参与方式有不同要求。

下表根据操作数的来源对比三种计算路径及其关键前提:

建议路径需求关键前提
WMMA从全局内存协作装载矩阵并执行 warp 级矩阵乘加A/B fragment 的 shape、类型和 layout 必须匹配。
SQMMA对已经搬入共享内存的数据执行较大矩阵乘加先完成 TME 搬运和线程同步,再创建 SQMMA descriptor。
CONV执行 direct convolutioninput、weight 和卷积参数必须来自匹配的 descriptor/编码流程。

请选择对应小节中的参数和调用约定。


WMMA

WMMA 用于在一个 warp 内协作执行矩阵乘加。它将 global memory 中的 A/B 矩阵 tile 装载到 fragment,在分布式寄存器中完成 D = A × B + C,再将 accumulator 写回 global memory。

当输入矩阵可以直接装载到 fragment,并且你需要显式控制 fragment 的 shape、数据类型和 layout 时,请选择 WMMA。参与同一操作的线程必须使用一致的控制流和参数组合;A/B fragment、accumulator 和 mma_syncM/N/K 必须匹配。

调用流程

你将声明 A/B fragment 和 accumulator,从 global memory 装载输入,执行 D = A × B + C,再将结果写回 global memory:

声明 A/B fragment 和 accumulator

从 global memory 装载 A 和 B

初始化 accumulator

执行矩阵乘加:D = A × B + C

将结果写回 global memory

下面的代码展示 WMMA 的基本调用顺序。WMMA API 是 warp 协作接口,参与同一操作的线程必须一致地创建和使用 fragment。

#include <musa_fp16.h>
#include <mma.h>
using namespace mtmusa;

__global__ void wmma_tile(
float *dst, const __half *a_ptr, const __half *b_ptr) {
wmma::fragment<wmma::matrix_a, 16, 16, 16,
__half, wmma::row_major> a;
wmma::fragment<wmma::matrix_b, 16, 16, 16,
__half, wmma::col_major> b;
wmma::fragment<wmma::accumulator, 16, 16, 16, float> c;

wmma::load_matrix_sync(a, a_ptr, 16);
wmma::load_matrix_sync(b, b_ptr, 16);
wmma::fill_fragment(c, 0.0f);
wmma::mma_sync(c, a, b, c, false);
wmma::store_matrix_sync(dst, c, 16, wmma::mem_row_major);
}

示例步骤

  1. 声明 A、B fragment 和 accumulator。
  2. 使用 load_matrix_sync 从 global memory 装载 A 和 B。
  3. 使用 fill_fragment 初始化 accumulator。
  4. 使用 mma_sync 执行矩阵乘加。
  5. 使用 store_matrix_sync 将 accumulator 写回 global memory。

matrix_amatrix_b fragment 需要 layout tag,accumulator 不需要。load_matrix_syncstore_matrix_syncldm 单位是元素,不是字节。A row/col 的典型 ldmK/M,B row/col 的典型 ldmN/K

使用一个 warp 启动:wmma_tile<<<grid, 32>>>(...)。WMMA 的 A/B fragment、mma_sync 和 accumulator 必须使用相同的 M/N/K

支持组合

请从下表选择匹配的 A/B 类型、累加类型、shape 和输出布局:

A / B 类型accumulator支持 shape(M x N x K)输出布局
__halffloat16x16x16, 16x16x32row/col
__mt_bfloat16float16x16x16, 16x16x32row/col
__mt_fp8_e4m3float16x16x16, 16x16x32row/col
__mt_fp8_e5m2float16x16x16, 16x16x32row/col
__half__mt_bfloat16signed charfloat16x16x16, 16x16x32row/col
signed charint16x16x16, 16x16x32, 16x16x64row/col
unsigned charint16x16x16, 16x16x32, 16x16x64row/col
wmma::precision::tf32float16x16x16row/col
wmma::precision::tf32float16x8x4row-major

TF32 输入存储为 float,使用对应的 TF32 fragment 类型;16x8x4 只采用 row-major 输出布局。

完整的函数签名、参数语义和支持边界请参见 WMMA API


SQMMA

SQMMA 用于对已经搬入 shared memory 的 A/B tile 执行矩阵乘加。它通过 descriptor 读取 shared memory 中的操作数,并将计算结果保存在 accumulator fragment 中。

当你需要显式组织 tile 搬运、共享内存布局和矩阵计算时,请选择 SQMMA。与直接从 global memory 装载 fragment 的 WMMA 不同,SQMMA 必须先通过 TME 完成 A/B tile 搬运,并使用 barrier 确认数据已经可读,再从 shared memory 创建与 shape、类型和 layout 匹配的 A/B descriptor。

调用流程

你将先准备 shared memory 中的 A/B tile,再创建 SQMMA descriptor,执行矩阵乘加并写回结果:

通过 TME 将 A/B tile 搬入 shared memory

等待 barrier,确认 tile 可读

创建 A/B SQMMA descriptor

初始化 accumulator

执行 SQMMA 矩阵乘加

将结果写回 global memory

当前示例按 128 线程组织:先将 tile 通过 TME 搬到对齐的共享内存,再构造 descriptor,最后执行 mma_sync。实际线程参与方式必须以目标 SDK 示例和接口约束为准。

下面的代码展示 SQMMA 的组织方式。它依赖具体的线程参与、descriptor、对齐和 shape 约束;请将其作为调用结构示例,并在使用目标 SDK 编译验证这些条件后再用于实际 kernel。

#include <musa_fp16.h>
#include <sqmma.h>
using namespace mtmusa;

template <typename T>
__global__ void sqmma_tile(float *dst, int *desc_a_addr, int *desc_b_addr) {
constexpr int M = 16, N = 64, K = 16;
__shared__ __align__(256) T a_smem[M * K];
__shared__ __align__(256) T b_smem[K * N];

__musa::async_barrier bar(1);
if (threadIdx.x < 32) bar.init_arrival(4);
__syncthreads();
if (threadIdx.x < 32)
__musa::memcpy_async(bar, a_smem, desc_a_addr,
M*K, 0, M*K*sizeof(T), 1, 3, 1);
if (threadIdx.x >= 32 && threadIdx.x < 64)
__musa::memcpy_async(bar, b_smem, desc_b_addr,
K*N, 0, K*N*sizeof(T), 2, 3, 1);
unsigned phase = bar.arrive();
bar.wait(phase);

sqmma::sqmmadesc<sqmma::desc_a, T> a_desc;
sqmma::sqmmadesc<sqmma::desc_b, T> b_desc;
sqmma::fragment<sqmma::accumulator, M, N, K, float> c;
sqmma::make_sqmma_desc(a_desc, a_smem, K*sizeof(T), sqmma::sg_16_byte);
sqmma::make_sqmma_desc(b_desc, b_smem, N*sizeof(T), sqmma::sg_32_byte);
sqmma::fill_fragment(c, 0.0f);
sqmma::mma_sync(c, a_desc, b_desc, c,
sqmma::mem_row_major, sqmma::mem_row_major,
sqmma::positive, sqmma::positive, sqmma::init_none);
sqmma::store_matrix_sync(dst, c, N, sqmma::mem_row_major);
}

示例步骤

  1. 通过 TME 将 A/B tile 搬入 shared memory,并等待 barrier。
  2. 使用 make_sqmma_desc 创建 A/B descriptor。
  3. 声明并初始化 accumulator fragment。
  4. 使用 mma_sync 执行 SQMMA 矩阵乘加。
  5. 使用 store_matrix_sync 将结果写回 global memory。

make_sqmma_desc 的 leading stride 单位是字节;sg_* 是共享内存 swizzle granularity。a_layout/b_layout 解释 descriptor 背后的逻辑布局,positive/negative 是输入 scale,init_zero/init_none 控制是否使用输入 C。

支持组合

FP16 / BF16 -> FP32

16x64x16 16x64x32 16x64x64
32x32x16 32x32x32 32x32x64 32x64x16 32x64x32 32x64x64
32x128x16 32x128x32 32x128x64
64x16x16 64x16x32 64x16x64 64x32x16 64x32x32 64x32x64
64x64x16 64x64x32 64x64x64 64x128x16 64x128x32 64x128x64
128x32x16 128x32x32 128x32x64 128x64x16 128x64x32 128x64x64
128x128x16 128x128x32 128x128x64

INT8 / UINT8 -> INT32

16x64x32 16x64x64 16x64x128
32x32x32 32x32x64 32x32x128 32x64x32 32x64x64 32x64x128
32x128x32 32x128x64 32x128x128
64x16x32 64x16x64 64x16x128 64x32x32 64x32x64 64x32x128
64x64x32 64x64x64 64x64x128 64x128x32 64x128x64 64x128x128
128x32x32 128x32x64 128x32x128 128x64x32 128x64x64 128x64x128
128x128x32 128x128x64 128x128x128

TF32 -> FP32

16x64x8 16x64x16 16x64x32 32x32x8 32x32x16 32x32x32
32x64x8 32x64x16 32x64x32 64x16x8 64x16x16 64x16x32
64x32x8 64x32x16 64x32x32 64x64x8 64x64x16 64x64x32
128x64x8 128x64x16 128x64x32 128x128x8 128x128x16 128x128x32

具体的 shape、数据类型、layout、scale 和 descriptor 参数请参见 SQMMA API


CONV

CONV 用于在 Device Code 中执行 direct convolution。它通过 descriptor 读取 shared memory 中的 input 和 weight tile,并结合 Host Code 编码的卷积参数完成计算。

当你需要显式控制卷积 tile、数据布局或融合逻辑时,请选择 CONV。计算前,必须先通过 TME 将 input 和 weight 搬入 shared memory,并使用 barrier 等待数据就绪;随后创建与实际内存布局匹配的卷积 descriptor,并使用对应的卷积参数执行计算。

调用流程

你将准备 input/weight tile 和 direct CONV 参数,创建 descriptor,执行卷积并写回结果:

通过 TME 将 input/weight tile 搬入 shared memory

准备 Host Code 编码的 direct CONV 参数

创建 CONV descriptor

初始化 accumulator

执行 direct convolution

将结果写回 global memory

下面的代码展示 CONV 的调用顺序。param_words、descriptor 类型和 shape 必须来自已验证的 Host/Device API 组合;完整参数约束仍以 API 参考为准。

conv::convdesc<conv::desc_a, T> a_desc;
conv::convdesc<conv::desc_b, T> b_desc;
conv::fragment<conv::accumulator, M, N, K, float> c;
conv::make_conv_desc(a_desc, a_smem);
conv::make_conv_desc(b_desc, b_smem);
conv::fill_fragment(c, 0.0f);
__musa::conv_param p(param_words[0], param_words[1], param_words[2],
param_words[3], param_words[4], param_words[5],
param_words[6], param_words[7]);
conv::conv_sync(c, a_desc, b_desc, c, p.get_conv_param(),
conv::positive, conv::positive, conv::init_none);
conv::store_matrix_sync(out, c, ldm, conv::mem_row_major);

示例步骤

  1. 通过 TME 将 input/weight tile 搬入 shared memory。
  2. 使用主机端 direct CONV 编码结果准备 __musa::conv_param
  3. 使用 conv::make_conv_desc 创建 input/weight descriptor。
  4. 初始化 accumulator,并使用 conv_sync 执行 direct convolution。
  5. 使用 conv::store_matrix_sync 将结果写回 global memory。

CONV 的 input/weight 先由 TME 搬到共享内存,再用 conv::make_conv_desc 包装。__musa::conv_param 必须来自主机端 direct CONV 参数编码,请直接使用编码结果构造参数。

支持组合

请从下表选择匹配的输入类型、累加类型和 shape:

输入 / weightaccumulator支持 shape(M x N x K)
__half__mt_bfloat16float128x64x64, 128x128x128, 256x16x16, 256x32x32, 256x64x64, 512x16x16, 512x32x32, 1024x16x16
conv::precision::tf32float同上 8 个 shape
signed charunsigned charint128x64x64, 128x128x128, 256x32x32, 256x64x64, 512x32x32

具体的卷积 descriptor、参数、shape 和数据类型约束请参见 CONV API


优化顺序

完成单 tile 的 Load → Compute → Store 后,再按以下顺序增加复杂度:

  1. barrier 和 transaction count;
  2. 多个 stage 和 pipeline;
  3. double buffering;
  4. TME prefetch;
  5. swizzle;
  6. im2col。

这些能力都可能改变线程分工、共享内存布局或参数组合。每增加一种能力,都应重新验证数据结果和同步顺序。


预取与 swizzle

LSU prefetch

如果需要提前提示 LSU 读取后续数据,可以使用以下预取接口:

prefetch(&x[1], 64);

这是 LSU 预取,不是 TME tensor copy。字节参数决定预取范围;代码应保证地址和范围有效。支持 int/float 的 64、128、256 字节形式。

TME prefetch

TME prefetch 没有独立主机 API,而是 __musa::memcpy_async overload 的尾参数。tile rank 1~5 和 im2col rank 3~5 采用各自对应的 overload,预取字节数必须与目标 cache line / 搬运粒度匹配。

TME swizzle

rank-1 tile load/store 可以使用以下 (SG, SS, SL) 组合:

(NONE,256B,256B)
(16B,32B,128B) (16B,64B,128B) (16B,128B,128B)
(32B,64B,128B) (32B,128B,128B)
(16B,32B,256B) (16B,64B,256B) (16B,128B,256B) (16B,256B,256B)
(32B,64B,256B) (32B,128B,256B) (32B,256B,256B)
(64B,128B,256B) (64B,256B,256B)

SG 表示 swizzle granularity,SS 表示 swizzle stride,SL 表示连续数据长度。三者必须与 tile 的连续字节数和共享内存布局相容。


编译

以下命令展示编译参数的组织方式。具体编译器路径、输入文件后缀和链接库名称必须在目标 SDK 环境中验证后再作为发布命令使用。

下面是使用 S5000 目标时的编译参数示例:

/usr/local/musa/bin/mcc -O2 -mtgpu --offload-arch=mp_31 \
demo.cu -o demo -lmusart -lmusa
./demo

正确性验证

在比较性能之前,请先确认:

  • descriptor 的地址、数据类型、rank、维度和 stride 正确;
  • tile shape、tile position 和 memcpy_async overload 使用同一 rank;
  • bytes、descriptor stride、leading stride 和 ldm 使用了正确的单位;
  • 所有参与 collective 操作的线程使用一致的 shape、类型、layout 和控制流;
  • 异步 TME load 的目标共享内存在对应 wait 之前不会被读取;
  • SQMMA 或 CONV 的共享内存操作数在创建 descriptor 前已经完成搬运和同步;
  • 不支持的 shape、alignment、rank 或目标架构保留普通路径作为回退。

只有这些条件都满足后,才适合比较 kernel-only 和端到端性能。


开发顺序

如果你第一次编写 Tensor kernel,建议按以下顺序推进:

  1. 验证 tensor descriptor;
  2. 验证单个 TME tile transfer;
  3. 加入 barrier,并验证 arrive/wait 顺序;
  4. 使用一个已确认支持的 WMMA、SQMMA 或 CONV 组合;
  5. 验证结果写回;
  6. 最后加入 pipeline、double buffering、prefetch、swizzle 或 im2col。

遇到编译或运行错误时,按 descriptor、rank、shape、alignment、字节数、barrier phase 和 stage 复用的顺序排查。每次只修改一个条件并重新验证,以便确定错误来源。


相关文档

更多详情,参见 TCE/TME API 参考协作组。前者介绍公开 API 的原型、类型选择器以及按 rank 区分的读写和预取语义,后者介绍通用线程组同步和异步搬运。