学习目标
学完本节你将能够:
- 理解 CUTLASS 如何通过多级流水线和共享内存分块实现高吞吐 GEMM
- 掌握
ldmatrix和mma.sync指令在 CUTLASS 中的角色 - 理解多阶段软件流水线的原理及其对隐藏内存延迟的作用
- 能够根据目标架构选择合理的 GEMM 参数(tile 大小、阶段数、warp 组织)
- 了解如何将 CUTLASS 集成到实际的深度学习算子中
1. CUTLASS GEMM 的整体数据流
CUTLASS 的 GEMM 实现遵循全局内存 → 共享内存 → 寄存器 → Tensor Core → 全局内存的数据流。其核心思想是通过多级缓存和流水线,最大化数据复用,最小化全局内存访问。
全局内存 (A, B)
│ ① 协作加载 (cp.async / 普通加载)
▼
共享内存 (Threadblock Tile: TileM×TileK, TileK×TileN)
│ ② ldmatrix / 普通加载
▼
寄存器 Fragment (每个线程持有 A/B 片段)
│ ③ mma.sync (Tensor Core 乘加)
▼
累加器 Fragment (每个线程持有 C 片段)
│ ④ 写回
▼
全局内存 (C)
关键优化:
- 使用异步拷贝(
cp.async)实现全局内存→共享内存的加载与计算重叠。 - 使用
ldmatrix指令高效地从共享内存加载矩阵片段到寄存器,并处理数据布局。 - 使用 Tensor Core 的
mma.sync指令执行矩阵乘加。 - 多阶段软件流水线让加载、计算、写回步骤重叠执行。
2. 共享内存分块与数据复用
2.1 Threadblock Tile 的选择
Tile 大小直接决定了共享内存占用和数据复用的程度。以 Ampere(sm_80)为例,常见配置:
TileM = 128, TileN = 128, TileK = 32(或 64)
共享内存占用计算公式:
- A tile:
TileM × TileK × sizeof(ElementA) - B tile:
TileK × TileN × sizeof(ElementB) - 若使用多级流水线(如 3 级),共享内存会根据阶段数成倍增加。
权衡关系
- 更大的
TileM/TileN:提升片上数据复用,但占用更多共享内存,会减少SM可驻留Block数量,降低占用率。 - 更大的
TileK:单次K维迭代计算量更大,但同样增加共享内存压力。
2.2 共享内存 Bank Conflict 处理
CUTLASS 在共享内存数组末尾自动插入 padding(填充),规避 Bank Conflict。
示例:逻辑上
[TileM][TileK],实际存储为[TileM][TileK+padding],改变访存索引映射关系,让不同Warp访问不同Bank。
2.3 多阶段流水线(Multi‑stage Pipeline)
CUTLASS 使用多个独立共享内存缓冲区(stage)实现流水线。例如 Stages = 3 代表3套缓冲区:
- Stage‑0:加载 tile k;
- Stage‑1:计算 tile k‑1;
- Stage‑2:写回 tile k‑2;
通过 cp.async + pipeline 提交/等待,加载新tile的同时计算上一个tile,掩盖全局内存访问延迟。K维度越大,流水线收益越明显。
代价:阶段数越大,共享内存开销越高,存在上限,受SM总共享内存约束。
3. 关键指令:ldmatrix 与 mma.sync
3.1 ldmatrix 指令
ldmatrix 是 Ampere(sm_80) 引入的硬件指令,从共享内存批量加载矩阵片段到寄存器。
- 支持加载8×8、16×8矩阵块;
- 硬件自动完成布局重排,适配Tensor Core的MMA指令;
- 可以做转置加载,适配B矩阵;
- 减少指令条数,规避部分Bank Conflict。
在CUTLASS内部,TensorOp模式下会自动使用
ldmatrix,配合TensorOpMultiplicand布局。
3.2 mma.sync 指令
mma.sync 是 Tensor Core 的底层硬件矩阵乘加指令,对应不同架构有不同的形状与精度:
| 架构 | 典型MMA指令 |
|---|---|
| Volta(sm70) | mma.sync.aligned.m16n8k8 |
| Ampere(sm80) | mma.sync.aligned.m16n8k16(FP16/BF16)、m16n8k8(TF32) |
| Hopper(sm90) | 支持FP8格式MMA |
CUTLASS通过InstructionShape模板参数,自动选择对应架构的MMA指令,用户一般不需要手写汇编。
4. CUTLASS 的代码生成与模板展开
CUTLASS 大量使用C++模板、#pragma unroll,全部tile、warp、stage参数均为编译期常量,编译器可以做激进优化:
- 内层循环完全展开;
- 消除边界分支判断;
- 优化寄存器分配;
- 直接生成贴近硬件SASS指令序列。
4.1 模板参数映射硬件资源
| 模板参数 | 对应硬件资源 |
|---|---|
ThreadblockShape |
Block级tile,决定共享内存大小、Block线程数 |
WarpShape |
Warp在tile内分工 |
InstructionShape |
Tensor Core MMA硬件指令形状 |
Stages |
共享内存流水线缓冲区数目 |
ElementA/B/C/Accumulator |
寄存器数量、计算精度 |
4.2 线程组织
Block内部线程被划分为多个Warp,每个Warp负责Warp‑tile。ThreadblockSwizzle控制Block在GPU上的排布顺序,优化L2缓存局部性。
5. 性能调优指南
5.1 不同架构参考配置
| 架构 | 推荐精度 | Tile大小 | Stages | 备注 |
|---|---|---|---|---|
| Volta (sm_70) | FP16 | 128×128×32 | 2‑3 | 第一代Tensor Core |
| Turing (sm_75) | FP16/INT8 | 128×128×32 | 2‑3 | 消费级Tensor Core,共享内存偏小 |
| Ampere (sm_80) | FP16/BF16/TF32 | 128×256×32 / 256×128×32 | 3‑4 | 支持cp.async,共享内存更大 |
| Hopper (sm_90) | FP8/FP16 | 更大tile | 4+ | 支持TMA分布式共享内存 |
5.2 调优步骤
- 确定精度:根据模型需求与硬件能力选择;
- 设置tile大小:从参考默认值起步,逐步调大,监控共享内存占用;
- 调整Stages:增大阶段数提升延迟隐藏,同时观察占用率是否下降;
- Nsight Compute分析:观测共享内存、寄存器、内存效率、Tensor Core利用率;
- 尝试WarpShape:修改Warp分块,观测Bank Conflict、SM利用率变化。
6. 代码演示:使用 CUTLASS 配置 GEMM
Ampere sm_80,FP16输入,FP32累加,TensorCore,3级流水线
#include <cutlass/cutlass.h>
#include <cutlass/gemm/device/gemm.h>
#include <cutlass/numeric_types.h>
#include <cutlass/util/host_tensor.h>
using ElementInput = cutlass::half_t;
using ElementOutput = float;
using ElementAccumulator = float;
using ThreadblockShape = cutlass::gemm::GemmShape<128, 128, 32>;
using WarpShape = cutlass::gemm::GemmShape<64, 64, 32>;
using InstructionShape = cutlass::gemm::GemmShape<16, 8, 16>;
using Gemm = cutlass::gemm::device::Gemm<
ElementInput,
cutlass::layout::RowMajor,
ElementInput,
cutlass::layout::ColumnMajor,
ElementOutput,
cutlass::layout::RowMajor,
ElementAccumulator,
cutlass::arch::OpClassTensorOp,
cutlass::arch::Sm80,
ThreadblockShape,
WarpShape,
InstructionShape,
3 // Stages 流水线阶段数
>;
int main() {
int M = 512, N = 512, K = 512;
cutlass::HostTensor<ElementInput, cutlass::layout::RowMajor> A({M, K});
cutlass::HostTensor<ElementInput, cutlass::layout::ColumnMajor> B({K, N});
cutlass::HostTensor<ElementOutput, cutlass::layout::RowMajor> C({M, N});
// 初始化A、B矩阵(省略)
cutlass::device_memory::allocation<uint8_t> workspace;
Gemm gemm_op;
cutlass::Status status = gemm_op({
{M, N, K},
{A.device_data(), K}, // A指针 + leading dimension
{B.device_data(), K}, // B指针 + leading dimension
{C.device_data(), N}, // C输入
{C.device_data(), N}, // D输出
{1.0f, 0.0f} // alpha, beta
}, nullptr, workspace);
// 同步、结果校验(省略)
return 0;
}
7. 课后练习
练习1:观察不同 Stages 的性能
修改 CUTLASS GEMM 示例中的 Stages 参数(1、2、3、4),测量执行时间,使用 ncu 观察共享内存占用和内存延迟隐藏效果。
练习2:调整 Tile 大小
将 ThreadblockShape 改为 256×128×32 或 128×256×32,比较性能,分析原因。
练习3:对比 TensorOp 与 Simt
分别配置 OpClassTensorOp 和 OpClassSimt 执行相同规模的 FP32 GEMM,对比性能和资源使用。
练习4:查看 SASS 中的 ldmatrix 和 mma
使用 nvdisasm 查看 CUTLASS GEMM 生成的 SASS,找到 LDSM 和 HMMA/MMA 指令,理解它们与代码的对应关系。
练习5:集成到自定义算子
将 CUTLASS GEMM 封装成一个函数,在你的项目中调用,并处理非方阵和边界情况。
8. 下一步
下一节将进入 CUTLASS 进阶:卷积、注意力与自定义算子,学习:
- 隐式GEMM卷积实现原理;
- FlashAttention基于CUTLASS的分块思想;
- Epilogue融合bias、ReLU等算子;
- 实战调优与常见坑点。