如何驯服 Blackwell 而不因 CUDA 发疯
说到编写 GPU 内核,我们大多数人会立刻想到无尽的 CUDA C++ 代码页、摆弄指针,以及调试内存泄漏或竞态条件的痛苦。当 NVIDIA Blackwell 架构(sm_100a)这样的庞然大物上市时,情况就更加复杂了。内部有太多的异步操作和新的数据引擎,旧的「只要把线程派到核心」的方法已经无法榨取出哪怕一半的性能了。
最近我偶然发现了 MLC-AI 团队的 modern-gpu-programming-for-mlsys 项目。从本质上讲,它不仅仅是一个代码仓库,更是一个关于现代 GPU 编程的活生生的交互式教程。这帮人全力以赴:他们教你使用 Python 和一种名为 TIRx 的新 DSL 来编写 SOTA(最先进)内核。
为什么这不是又一个普通教程
通常 GPU 课程都会卡在 2010 年代初期的经典矩阵乘法上。而这里,焦点转向了明天的硬件。Blackwell 架构、其内存层次结构、张量核心和异步数据移动机制是关注的中心。
作者们提出了一条从理解硬件到编写在现代 ML 系统中实际运行的内核的路径。你不需要与底层 C++ 搏斗,而是用 TIRx(Tensor IR next)编写。这是一个 Apache TVM 扩展,让你在 Python 中描述内核逻辑,同时保持对寄存器、缓存和张量引擎的控制。
这个教程里有什么
所有内容被分解为逐步提升难度的逻辑块。
首先,他们讲解执行模型和内存模型。这里没有枯燥的定义,而是大量关于 TMA(Tensor Memory Accelerator)的实践工作,以及理解数据如何真正到达计算单元。有趣的是,作者们详细阐述了 Roofline 模型——它有助于理解你的代码受什么瓶颈限制:是内存带宽还是核心的原始算力。
然后是对 TIRx 的深入探讨。你学习使用 tiles(TileLayout)、axes 和 swizzle 机制(在内存中进行巧妙的数据混排以避免 bank 冲突)。最后,你将得到一个可用的 GEMM(矩阵乘法)内核。
之后就是真正的「硬菜」。作者展示了如何将普通的 GEMM 变成高性能怪兽。TMA 流水线、warp 特化化以及将 CTA 分组到集群中都会派上用场。如果这些术语听起来像咒语,教程会巧妙地将它们一一解释清楚。
在最终环节,你将创建 Flash Attention 4。这正是现代大语言模型运行的内核。该内核包含两个 MMA(矩阵-矩阵累加)阶段、在线 softmax 和 GQA(分组查询注意力)支持。
技术栈和特性
该项目与 Apache TVM 生态系统紧密耦合。要运行示例,你需要:
- 库
apache-tvm(TIRx 包含在其中)。 - PyTorch 用于验证结果和准备数据。
- 一块 Blackwell GPU(例如 B200)如果你想在真实硬件上运行代码。
如果你手边没有 Blackwell(这对我们大多数人来说目前是合理的),这个仓库作为理论基础仍然极其有用。内存优化和计算调度的逻辑也适用于较旧的架构,尽管某些特性如 TMA 是特定于较新的 NVIDIA 系列的。
要在本地构建这本书并以方便的格式阅读,标准的 Sphinx 就足够了:
pip install -r requirements-docs.txt
sphinx-build -b html . _build/html
这值得你花时间吗
我经常遇到一些 ML 基础设施开发者,他们害怕深入内核编写,认为这是「C++ 大神」的专利。这个项目证明了,有了正确的抽象(如 TIRx)和对架构的理解,你可以在 Python 生态系统中编写世界级的代码。
这个项目对于那些致力于在「底层」优化神经网络或想理解为什么 Flash Attention 比常规方法更快的人特别有用。即使你明天不打算发布自己的 Blackwell 库,深入了解张量内存和异步队列的工作原理确实能开阔思路,改变你对高性能计算的看法。
唯一的缺点是入门门槛仍然很高。这不是「给傻瓜的 GPU」,而是「给想成为专家的 GPU」。但如果你准备好弄清楚字节在显卡内部究竟是如何移动的,目前很难找到比这更好的资料了。
相关项目