NVIDIA Blackwell 架构对 GEMM 核心实现带来了革新。本文旨在深入探讨其核心特性 UMMA 指令与 Tensor Memory (TMEM),通过解析最小工作示例,阐明如何更新或从头编写高性能的 Blackwell GEMM 核心,为开发者提供清晰的实践路径与技术要点。
智能速览
Blackwell 架构弃用 Hopper 的 WGMMA,转而采用 UMMA 指令进行矩阵运算。
Tensor Memory (TMEM) 是专供张量核心使用的片上内存,可显著降低寄存器压力。
UMMA 指令支持单线程异步执行,实现了 MMA 操作与 CTA 主流程的深度解耦。
TMEM 采用二维编址方式,通过 tcgen05.alloc 指令动态分配,并由同一 warp 释放。
CUTLASS/CuTe 提供了新的原子与布局抽象,简化了 UMMA 和 TMEM 的使用。
精华内容
理解 Blackwell 的革新,关键在于掌握其新指令与专用内存。下面将通过一个最小示例,拆解 UMMA 与 TMEM 的具体应用方法。
UMMA 与 TMEM 概览
Blackwell 架构引入了 UMMA(即 tcgen05.mma)指令,用以取代 Hopper 的 WGMMA,为张量核心提供计算能力。与 WGMMA 不同,UMMA 期望其输入或累加器位于一种新型的片上内存——Tensor Memory (TMEM) 中。
TMEM 是每个 SM 大小为 256KB 的专用存储资源,采用二维组织结构(512列x128通道)。它的主要目的是替代通用寄存器,直接为张量核心服务。此举将 MMA 操作从寄存器中解放出来,大幅降低了寄存器压力,并使得整个操作可以由单个线程异步启动,进一步与 CTA 的主执行流程解耦。
PTX 指令详解
在 PTX 层面,tcgen05.mma 指令定义了矩阵运算的细节,其支持的最大原子形状可达 128x256x16,是 Hopper 最大 WGMMA 原子的两倍。该指令需要操作数描述符(描述 SMEM 中的数据布局)和一个指令描述符(包含数据类型、转置等信息)。
计算完成后,结果存于 TMEM 中,必须使用 tcgen05.ld 指令将其加载到寄存器以进行后处理。tcgen05.ld 是一条 warp 级同步指令,其访问模式受限,每个 warp 只能访问 TMEM 的 32 个通道,这要求数据移动操作必须由整个 warp 组协同完成。
CUTLASS 中的实现
CUTLASS 3.x 中的 CuTe 库为 Blackwell 的新特性提供了高级抽象。新的 MMA 原子(如 SM100_MMA_F16BF16_SS)通过模板参数直接映射到 UMMA 指令的特性,包括操作数数据类型、矩阵形状以及是否对输入进行转置或取反。
一个显著变化是,由于 UMMA 是单线程指令,其内部的“线程布局”概念被重新用于表示协同执行该指令的 CTA 布局。在单 CTA 示例中,这些布局的维度常为 1,简化了代码结构,也为后续多 CTA 协作的设计奠定了基础。
TMEM 的管理与同步
TMEM 的使用涉及严格的分配与释放生命周期。开发者需使用 cute::TMEM::Allocator1Sm 辅助类,它封装了 tcgen05.alloc 和 tcgen05.dealloc 指令。分配和释放操作必须由同一个 warp 执行,且分配的列数需满足特定约束(如 2 的幂)。
由于 UMMA 是异步的,其执行需要同步机制。示例代码中使用了 mbarrier(内存屏障),其工作流程与 TMA 的同步类似,由负责启动 UMMA 的 warp 中的一个线程进行初始化,确保 MMA 操作完成后再进行后续步骤。
数据后处理流程
当 UMMA 计算结束后,累加器矩阵仍在 TMEM 中。为了进行后续的元素级操作(如激活函数)或写回全局内存,必须将数据移至寄存器。这一过程通过 CUTLASS 抽象的拷贝原子(如 SM100_TMEM_LOAD_32dp32b1x)来完成。
该原子对应 tcgen05.ld 指令,它规定了数据移动的形状。由于每个 warp 只能访问 32 个 TMEM 通道,一个包含 4 个 warp 的 warp 组需要协同工作,才能完整地从 TMEM 中提取出 128x256 大小的结果块,每个线程负责加载数据的一个 32 位单元。
通过此文,我们掌握了 Blackwell 架构下 GEMM 核心的关键变革,特别是 UMMA 与 TMEM 的协同工作原理。这些优化显著提升了资源利用效率。未来,我们将进一步探索多 SM 协作与更复杂集群模式下的实现策略。