Triton Gluon 中的 AMD RDNA4 WMMA API:面向 gfx1200/gfx1201 的矩阵乘内建指令解析
【免费下载链接】tritonDevelopment repository for the Triton language and compiler项目地址: https://gitcode.com/GitHub_Trending/tri/triton
Gluon 是 Triton 面向高级内核的底层 GPU 编程模型,将布局、共享内存、warp specialization 与目标特性直接暴露给开发者,用"牺牲便利换取控制力"的方式支撑极致性能内核。本文以 docs/gluon/api/amd.rdna4.rst 文档页为核心,深入讲解其中唯一的公开 API——triton.experimental.gluon.language.amd.rdna4.wmma:它如何在 RDNA4(gfx1200/gfx1201)上驱动 AMD WMMA 矩阵乘指令,其参数语义、底层布局约束与完整调用链是什么。读完本文,你将能够在自己的 Gluon 内核中正确使用 WMMA 内建函数,并理解它与其他 AMD 代际 API(RDNA3、CDNA 系列)之间的关系。
一、文档定位:RDNA4 代际 API 参考页
在 Gluon 的 AMD API 体系(docs/gluon/api/amd.rst)中,AMD 目标相关的 Gluon API 按 GPU 代际拆分为多个独立参考页:CDNA 3、CDNA 4、CDNA 5、RDNA 3、RDNA 4。其中 RDNA 4 参考页 amd.rdna4.rst 通过autosummary机制将triton.experimental.gluon.language.amd.rdna4模块中带文档字符串的公开符号自动收集并生成 API 文档,该模块导出的全部符号只有一个:
wmma——使用 AMD WMMA 指令完成a * b + acc矩阵乘的内建函数。
与 CDNA 4 参考页 amd.cdna4.rst 中丰富的mfma、mfma_scaled、async_copy、buffer_atomic_*等十余个符号相比,RDNA4 页面的 API 面极其精简。这不是文档遗漏,而是 RDNA 架构的真实反映:RDNA 代际面向图形与消费级计算,矩阵乘硬件路径集中于 WMMA(Wave Matrix Multiply-Accumulate)指令,而 CDNA 代际面向数据中心 AI 计算,提供 MFMA 矩阵核以及完整的缩放、异步拷贝与原子操作集合。从 rdna4/init.py 与 cdna4/init.py 两个模块的导出列表对比中,可以清晰看到这一代际分工。
二、wmma 内建函数:签名与语义
RDNA4 的wmma定义于 python/triton/experimental/gluon/language/amd/rdna4/init.py:
from ..._core import builtin from .._ops import _wmma __all__ = ["wmma"] @builtin def wmma(a, b, acc, _semantic=None): """ Computes matrix multiplication ``a * b + acc`` using an AMD WMMA instruction. Args: a (tensor): The operand a to be multiplied. b (tensor): The operand b to be multiplied. acc (tensor): The accumulator tensor. """ return _wmma(2, a, b, acc, _semantic)关键点:
- 签名:
wmma(a, b, acc)三个张量参数,含义即文档字符串所述——计算a * b + acc。_semantic为内部参数,由 Gluon 前端注入,普通用户无需关心。 - 代际版本号:模块把
_wmma的version固定为2。对照 amd/_layouts.py 中AMDWMMALayout的版本注释:版本1对应 RDNA3(gfx1100、gfx1101),版本2对应RDNA4(gfx1200、gfx1201),版本3对应 CDNA5(gfx1250)。因此同一份_wmma共享实现被 RDNA3 与 RDNA4 两个模块以不同版本号复用(对比 rdna3/init.py 中调用_wmma(1, ...))。 - 返回类型:返回与
acc相同类型(ttgl.tensor(handle, acc.type))的结果张量,即累加器布局同时决定结果的分布式类型。
三、共享实现_wmma:校验、降级与调用链
wmma的实质逻辑位于 python/triton/experimental/gluon/language/amd/_ops.py:
def _wmma(version, a, b, acc, semantic): """ Shared implementation for AMD WMMA operations for Gluon builtins """ _verify_wmma(version, a, b, acc) handle = semantic.dot(a, b, acc, input_precision=knobs.language.fp32_default, max_num_imprecise_acc=None, out_dtype=acc.dtype).handle return ttgl.tensor(handle, acc.type)调用链可以概括为:wmma(模块内建)→_wmma(共享实现)→semantic.dot(Gluon 语义层的 dot 内建)→ 后端指令选择与代码生成。_wmma不做任何数值上的自定义运算,而是委托给 Gluon 的通用矩阵乘语义semantic.dot,并显式传入:
input_precision=knobs.language.fp32_default:输入精度取自 Triton 的全局 knoblanguage.fp32_default(见 python/triton/knobs.py),即 FP32 输入的默认计算精度策略;max_num_imprecise_acc=None:不限制不精确累加次数;out_dtype=acc.dtype:输出类型跟随累加器。
也就是说,WMMA 的硬件指令选择是由后端根据AMDWMMALayout布局自动完成的,Gluon 前端只负责把语义合法的a * b + acc交给语义层,这一点与 CDNA 的mfma内建(直接绑定 MFMA 布局与缩放语义)存在架构上的差异——RDNA4 的wmma更接近"带布局约束的 dot"。
3.1 布局校验_verify_wmma
调用semantic.dot之前,_verify_wmma(amd/_ops.py)会对三个操作数做严格的布局约束检查,任何一条不满足都会抛出断言错误:
acc必填,且其布局必须是AMDWMMALayout,并且layout.version == version(RDNA4 即版本 2);a的布局必须是DotOperandLayout,且其parent是与acc版本一致的AMDWMMALayout;b的布局同样必须是DotOperandLayout,parent 为匹配的AMDWMMALayout。
这印证了 Gluon 的布局体系:累加器持有一个"父级" WMMA 布局,两个乘数操作数则通过DotOperandLayout的parent字段挂靠到该父布局上,从而把 M/N/K 维的线程与寄存器映射关系统一起来。若你在内核中手写wmma调用而操作数布局不满足上述关系,会在编译期收到明确的断言错误。
四、AMDWMMALayout:RDNA4 WMMA 的布局载体
理解wmma必须理解其布局类型AMDWMMALayout,定义于 python/triton/experimental/gluon/language/amd/_layouts.py。它继承自 Gluon 核心的DistributedLayout(python/triton/experimental/gluon/language/_layouts.py),核心字段如下:
| 字段 | 类型 | 含义 | 默认值 |
|---|---|---|---|
version | int | GPU 架构代际,1=RDNA3,2=RDNA4,3=CDNA5 | 必填 |
transposed | bool | 结果张量是否转置(影响线程持有连续元素的排布) | 必填 |
warp_bases | List[List[int]] | CTA 布局的 warp 基向量 | 必填 |
reg_bases | Optional[List[List[int]]] | CTA 布局的重复(寄存器)基向量 | [] |
instr_shape | Optional[List[int]] | 指令形状 (M, N, K) | [16, 16, 16] |
cga_layout | List[List[int]] | CTA 平铺(cluster)基向量 | [] |
rank | Optional[int] | warp/寄存器基的秩 | 2 |
值得注意的实现细节(amd/_layouts.py):
instr_shape缺省为[16, 16, 16],即单个 WMMA 指令的 (M, N, K) 形状;RDNA 的 WMMA 指令族正是以 16×16×16 为典型基本形状。rank缺省为 2,对应二维的 warp/寄存器平铺;cga_layout为空时表示不做 CTA 级平铺。- 所有字段在
__post_init__中先经_unwrap_if_constexpr解包(允许 constexpr 传参),再调用verify()做合法性校验。 _to_ir方法把布局序列化为 IR 层的get_amd_wmma_layout调用,这是 Gluon 前端与 MLIR 后端之间的桥梁。
同时,AMDMFMALayout(同文件 L16-L106)是 CDNA 系列的对应布局,其verify()会检查instr_shape的前两维属于[[32,32], [16,16], [64,4], [4,64]]、element_bitwidth为 32 或 64,版本区间为 1(gfx908)到 4(gfx950)——与 WMMA 布局的代际映射(1=gfx1100/1101,2=gfx1200/1201,3=gfx1250)共同构成 AMD 各代 GPU 的布局版本表。从源码结构看,MFMA 布局覆盖了从 CDNA1 到 CDNA4 的更长时间跨度,而 WMMA 布局覆盖 RDNA3/RDNA4 与 CDNA5,两条指令族在 CDNA5 上汇合(AMDWMMALayout版本 3)。
五、在内核中使用 RDNA4 wmma
结合上述签名与布局约束,一个典型的用法模式是:先为累加器构造AMDWMMALayout(version=2, ...),再让a、b以DotOperandLayout挂靠该布局,最后调用wmma(a, b, acc)。例如:
from triton.experimental.gluon import language as ttgl from triton.experimental.gluon.language.amd.rdna4 import wmma from triton.experimental.gluon.language.amd._layouts import AMDWMMALayout from triton.experimental.gluon.language._layouts import DotOperandLayout # RDNA4 (gfx1200/gfx1201) 上的 WMMA 累加器布局 acc_layout = AMDWMMALayout( version=2, # RDNA4 transposed=False, warp_bases=[[4, 1], [1, 4]], instr_shape=[16, 16, 16], ) # 操作数布局:DotOperandLayout 的 parent 必须指向 acc_layout a_layout = DotOperandLayout(parent=acc_layout, operand_index=0, k_width=16) b_layout = DotOperandLayout(parent=acc_layout, operand_index=1, k_width=16) # ... 构造 a、b、acc 三个分布式张量 ... c = wmma(a, b, acc) # c = a * b + acc需要强调:
- 上述布局参数仅为示意,实际内核中布局通常由后端推导(例如通过
ttgl.make_tensor_descriptor或自动布局推导流程获得),手动构造布局属于 Gluon 提供的"高级控制"路径; - 若在 RDNA3(gfx1100/gfx1101)上运行,应改用 amd.rdna3 模块 的
wmma(内部版本号为 1),两者 API 签名完全一致; - 完整可运行的 Gluon 示例内核参见 python/tutorials/gluon/ 教程目录与 python/examples/gluon/ 示例目录,Gluon 总览 docs/gluon/index.rst 也提供了教程与示例画廊的入口。
六、与相邻代际 API 的对照
为帮助你在正确的硬件上选择正确的 API,下表总结了当前仓库中 AMD 各代际 Gluon 模块的实际导出情况(依据各__init__.py源码):
| 代际 | 模块路径 | 代表性内建 | 布局版本 |
|---|---|---|---|
| RDNA 3 | language/amd/rdna3/ | wmma | AMDWMMALayoutv1(gfx1100/1101) |
| RDNA 4 | language/amd/rdna4/ | wmma | AMDWMMALayoutv2(gfx1200/1201) |
| CDNA 3 | language/amd/cdna3/ | buffer_*、mfma等基础集 | AMDMFMALayoutv3(gfx942) |
| CDNA 4 | language/amd/cdna4/ | mfma_scaled、scaled_upcast/downcast、load_shared_fp4_repacked、async_copy、buffer_atomic_* | AMDMFMALayoutv4(gfx950) |
| CDNA 5 | language/amd/cdna5/ | 扩展集(含AMDWMMALayoutv3) | MFMA + WMMA 并存 |
从实现上看(cdna4/init.py),CDNA4 模块会通过from ..cdna3 import *继承 CDNA3 的全部符号,再叠加 MX 格式(Microscaling)相关的缩放矩阵乘能力;而 RDNA4 模块仅导出wmma一个符号,专注于 WMMA 路径。选择哪个模块,取决于目标硬件属于 RDNA 消费级架构还是 CDNA 数据中心架构。
七、结语
docs/gluon/api/amd.rdna4.rst虽然只有寥寥数行的autosummary指令,但它指向的rdna4.wmma内建函数承载了完整的代际语义:版本号2精确映射 gfx1200/gfx1201,共享实现_wmma通过_verify_wmma强制AMDWMMALayout+DotOperandLayout的父子布局结构,并委托semantic.dot完成最终指令生成。对 Gluon 开发者而言,掌握这一条调用链与布局约束,就等于掌握了 RDNA4 上"布局即性能"的编程范式核心。
参考资源(仓库内)
- API 参考页:docs/gluon/api/amd.rdna4.rst、docs/gluon/api/amd.rst、docs/gluon/api/amd.rdna3.rst、docs/gluon/api/amd.cdna4.rst
- 内建实现:rdna4/init.py、rdna3/init.py、amd/_ops.py
- 布局定义:amd/_layouts.py
- Gluon 总览:docs/gluon/index.rst
- 教程与示例:python/tutorials/gluon/、python/examples/gluon/
【免费下载链接】tritonDevelopment repository for the Triton language and compiler项目地址: https://gitcode.com/GitHub_Trending/tri/triton
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考