CUTLASS Operator API 参数体系完全解析:RuntimeArguments、Operands 与类型标记
【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass
导读
CUTLASS Operator API 是 NVIDIA CUTLASS 提供的 Python 集成层,用于统一管理用 CuTe DSL 等 Python DSL 编写的高性能线性代数 kernel。本文围绕官方 API 参考页 arguments.rst 中定义的Arguments and Operands(参数与操作数)体系,深入讲解RuntimeArguments、GemmArguments、GroupedGemmArguments、EpilogueArguments、DenseTensor/ScaledOperand以及TensorLike/NumericLike类型标记的设计与使用。读完本文,你将掌握如何构造一次 CUTLASS 算子调用所需的完整参数对象(含块缩放 GEMM 与自定义 epilogue 融合),并理解参数从"用户友好对象"到"内部可编译表示"(TensorWrapper与cutlass.Numeric)的自动转换机制。
一、总体设计:一次算子调用 = RuntimeArguments + Operands
在 CUTLASS Operator API 中,一次操作(operation)被表达为一个RuntimeArguments对象,它同时描述:
- 操作类型(operation kind):例如 GEMM、Grouped GEMM;
- 操作数(operands):该操作实际运算的张量。
每种操作类型都有自己专属的RuntimeArguments子类(如GemmArguments、GroupedGemmArguments),实现同一操作类型的所有 Operator 接受同一个RuntimeArguments子类,这就是"kernel 无关接口"的基石。操作数则由Operand子类描述,例如封装单个稠密张量的DenseTensor、封装"量化张量 + 缩放因子张量"的ScaledOperand。
参数对象内部,张量类型字段与数值类型字段分别接受满足TensorLike/NumericLike协议的任何对象:
TensorLike:torch.Tensor、cute.Tensor、cutlass.operators.utils.tensor.TensorWrapper等;NumericLike:cutlass.Numeric、torch.dtype等。
该设计意味着你可以直接用 PyTorch 张量(或其他 DLPack 兼容张量)构造参数,而无需为不同 kernel 编写框架张量转换胶水代码——这正对应 operators/README.md 中"pass PyTorch tensors directly intoGemmArgumentsand calloperator.run(args)"的核心理念。所有公开符号统一从 cutlass/operators/init.py 导出(RuntimeArguments、GemmArguments、EpilogueArguments、DenseTensor、ScaledOperand、ScaleMode、ScaleSwizzleMode、TensorLike、NumericLike等),通常以import cutlass.operators as ops方式使用。
二、RuntimeArguments:参数基类与运行时性能控制
2.1 基类定义
RuntimeArguments是定义在 operators/cutlass/operators/arguments/base.py 中的抽象 dataclass:
@dataclass class RuntimeArguments: """Describes the operands and all other arguments passed to an Operator at runtime. ...""" performance: PerformanceControls | None = field(default=None, kw_only=True) """Optional runtime performance controls passed to the Operator""" def _validate(self): """Checks that the arguments are valid. This is run before all fields have been converted to TensorWrapper and cutlass.Numeric.""" def __post_init__(self): _convert_to_internal_types(self)从源码结构看,基类承担两个职责:
- 携带运行时性能控制:可选的
performance字段(kw_only=True,关键字专用),指向PerformanceControls实例; - 触发内部类型转换:
__post_init__中调用_convert_to_internal_types(self),把用户传入的框架张量、torch.dtype等统一转换为内部表示。
2.2 PerformanceControls
PerformanceControls(同文件 base.py)是"所有运行时性能选项的通用容器",本身不强制任何字段——不同的算子/实现可以为其定义具体的子类字段(如 tile 调度、工作区等),从而在不改变参数接口的前提下支持运行时性能调优。
2.3 类型自动转换机制(核心原理)
_convert_to_internal_types(base.py)是整套参数体系的"转换引擎"。它利用 dataclass 的类型注解(get_type_hints)逐字段检查:
| 字段注解 | 转换目标 | 说明 |
|---|---|---|
TensorLike | TensorWrapper | 用TensorWrapper(value, **global_metadata)包装,负责编译期/运行期张量描述(见下文) |
NumericLike | cutlass.Numeric | 经cutlass.operators.utils.dtype.to_cutlass_type转换为 CUTLASS 数值类型(支持torch.dtype等) |
实现了_convert_to_internal_types的对象 | 递归转换其内部字段 | 例如ScaledOperand、EpilogueArguments等复合对象 |
已经是TensorWrapper | 原样保留 | 避免重复包装 |
关键点:TensorWrapper是 operators/cutlass/operators/utils/tensor.py 中定义的"双张量"包装器,同时持有:
runtime_tensor:运行期使用的张量(真实数据);compile_time_tensor:编译期使用的张量(TVM-FFI 关闭时直接用cute.Tensor;开启时使用 fake tensor)。
这样无论是否启用 TVM-FFI,上层接口保持不变。此外TensorWrapper还处理了亚字节打包 dtype(如float4_e2m1fn_x2每个字节存 2 个 FP4 值):构造时会把物理 shape/stride 展开为逻辑 shape/stride,保证逻辑布局一致。
三、GemmArguments:GEMM 操作参数
3.1 字段与便捷构造
GemmArguments定义在 operators/cutlass/operators/arguments/gemm.py,表示out = A @ B的 GEMM 运算:
| 字段 | 类型 | 含义 |
|---|---|---|
A | Operand | 输入张量 A,shape(L, M, K)或(M, K) |
B | Operand | 输入张量 B,shape(L, K, N)或(K, N) |
out | Operand | 输出张量 C,shape(L, M, N)或(M, N) |
accumulator_type | NumericLike | 累加器数据类型 |
epilogue | EpilogueArguments \| None | 可选的 GEMM 后自定义 epilogue 融合 |
performance | PerformanceControls \| None | 继承自基类的运行时性能控制(关键字参数) |
其中M= A 与 out 的行数,N= B 与 out 的列数,K= A 的列数与 B 的行数,L= 批量矩阵乘法批数。所有张量必须同为 rank-3 或同为 rank-2。
便捷构造:对稠密 GEMM,A、B、out可以直接传裸张量而不必包DenseTensor:
GemmArguments(A, B, out, accumulator_type) # 等价于: GemmArguments(DenseTensor(A), DenseTensor(B), DenseTensor(out), accumulator_type)而其他操作数类型必须显式包装。例如块缩放 GEMM(scaled GEMM):
GemmArguments( ScaledOperand(A, ScaleATensor, scale_mode, scale_swizzle), ScaledOperand(B, ScaleBTensor, scale_mode, scale_swizzle), out, # 无需 DenseTensor 包装 accumulator_type, )该便捷逻辑由_operand_or_dense(operand.py)实现:若传入对象已是Operand则直接返回,若是TensorLike则包装为DenseTensor,否则抛出ValueError。
3.2 问题尺寸派生:GemmProblemSize
GemmProblemSize(gemm.py)是一个NamedTuple,其(M, N, K, L)与操作数形状对应关系为:
A:(L, M, K)B:(L, K, N)out:(L, M, N)
GemmArguments.problem_size属性直接由操作数形状推导(gemm.py):M = A.shape[-2]、N = B.shape[-1]、K = A.shape[-1]、L = A.shape[0] if rank==3 else 1。这意味着算子发现(ops.get_operators)无需用户显式声明问题尺寸,全部从张量自描述。
3.3 参数校验(_validate)
构造时GemmArguments会调用_validate()(gemm.py),执行以下一致性检查,不通过即抛出带完整形状信息的ValueError:
- A、B、out 必须是 rank 2 或 3;
- A 的 K 维与 B 的 K 维相等(注意
_logical_contraction_extent会考虑亚字节 dtype 的 K 维打包系数,见 gemm.py); - out 的 M 维等于 A 的 M 维;
- out 的 N 维等于 B 的 N 维;
- A、B、out 的 batch 维(如有)必须一致。
构造顺序(__post_init__,gemm.py):先_validate(),再_convert_epilogue()(用累加器形状(L, M, N)与类型对 epilogue 做 trace 并转换为TensorWrapper),最后调用基类__post_init__完成内部类型转换。
四、GroupedGemmArguments:分组 GEMM 参数
GroupedGemmArguments定义在 operators/cutlass/operators/arguments/grouped_gemm.py,执行一系列各自维度可以不同的独立 GEMM。抽象形式为:
for i in range(problems_in_group): out[i] = A[i] @ B[i]实际使用中,各问题的张量常被连续拼接到单个大张量里,此时需要offsets张量标出组内每个问题的结束位置。当前仓库支持的是Contiguous offset 2D-3D grouped GEMM变体:
A:shape(TotalM, K)或(1, TotalM, K)B:shape(problems_in_group, K, N)out:shape(TotalM, N)或(1, TotalM, N)offsets:标出每个问题结束位置的张量
start = 0 for i in range(problems_in_group): end = offsets[i] out[start:end, :] = A[start:end, :] @ B[i, :, :] start = end其字段在GemmArguments基础上增加了offsets(Operand类型,dataclass 字段带有metadata={"alignment_bytes": 4},即要求 4 字节对齐,供TensorWrapper使用)。构造函数签名(grouped_gemm.py):
GroupedGemmArguments( A, B, out, accumulator_type, offsets=offsets, epilogue=None, )集成测试 operators/test/integration/test_contiguous_offset_dense_gemm.py 演示了完整流程:构造offsets_real = torch.Tensor([128, 256])(int32、每个问题结束位置),随后operator.run(args)并以torch._grouped_mm的结果作为参考比对。注意测试还展示了错误用法会被算子发现机制拒绝:offsets元素个数不等于problem_count时,ops.get_operators返回空列表(test_contiguous_offset_dense_gemm.py)。
五、EpilogueArguments:自定义 epilogue 融合参数
5.1 设计目标与 epilogue_fn 约束
EpilogueArguments(operators/cutlass/operators/arguments/epilogue.py)描述一个融合在主操作之上的用户自定义 epilogue:它接收矩阵乘法结果,做张量级变换后写回输出。其定义是"泛型"的——接受一个epilogue_fn加任意**kwargs,底层通过对epilogue_fn的AST 解析确定输入输出。
epilogue_fn必须满足以下约束:
- 第一个位置参数必须命名为
accum——主操作(如 GEMM 的A @ B)的结果; - 至少返回一个张量,且返回列表中必须有一个名为
D的输出; accum之后的每个参数都是要加载的张量/标量;- return 语句中的每个变量都是要存储的张量/标量;
- 函数体必须满足**静态单赋值(SSA)**形式——每个变量只能被赋值一次。
函数一般结构:
def custom_epi_name(accum, *args) -> TensorType | tuple[TensorType, ...]: # Do some compute return D # and potentially other values5.2 kwargs:样例张量即"规格说明"
kwargs必须为 epilogue 中出现的所有输入与输出提供样例张量/标量。例如对于:
def my_epi(accum, alpha, C, beta): F = (accum * alpha) + (C * beta) D = relu(F) return D, F需要构造:
epi_args = EpilogueArguments( my_epi, alpha=..., C=..., beta=..., D=..., F=... )构造函数内部(epilogue.py)先调用trace_in_out(epilogue_fn)解析出输入/输出参数名列表,再从kwargs提取对应张量存入有序字典tensors(同名去重,因为一个名字可同时是输入与输出,如def epi(accum, D): ... return D);若有未在 AST 中出现的多余 kwarg 则抛ValueError。parameters/parameter_names属性返回参数值与名称列表。
5.3 trace:生成内部表示
trace(accumulator_shape, accumulator_type)(epilogue.py)以EmptyTensor构造累加器占位(LayoutType.RowMajor),连同self.tensors一起交给cutlass.operators.fusion.trace解析 AST、生成内部表示(DAG IR),并做有限的正确性检查(如形状匹配)。随后to_tensor_wrappers()将参数转为TensorWrapper;其中标量归约输出(如total = sum(accum)这类ScalarReductionImpl目的地)会被强制标记为static_layout(epilogue.py),以保证 EFC kernel 的跨 CTA atomic 归约路径可用。
5.4 Load / Store / Transport:逐张量数据搬运策略
默认情况下,epilogue 的张量通过TMA搬运。EpilogueArguments支持用Load/Store描述符包装 kwargs 来覆盖默认搬运方式(epilogue.py):
C=ops.Load(C, via=ops.Transport.ASYNC_GMEM_LOAD) D=ops.Store(D, via=ops.Transport.SYNC_GMEM_STORE)Transport枚举(epilogue.py)镜像 EFC kernel 的搬运目录:
| 枚举值 | 含义 |
|---|---|
TMA | 经共享内存的 TMA 搬运(默认) |
SYNC_GMEM_LOAD/SYNC_GMEM_STORE | 直接 GMEM 寻址,对发起线程同步 |
ASYNC_GMEM_LOAD | 经cp.async异步经共享内存读入 |
约束:Load 仅允许TMA/SYNC_GMEM_LOAD/ASYNC_GMEM_LOAD;Store 仅允许TMA/SYNC_GMEM_STORE。可选num_bits_per_copy指定非 TMA 传输的事务位宽(None时自动推导);它必须是 int 且不能与 TMA 组合使用(会在__post_init__中早期报错,避免静默丢弃)。
5.5 测试实证
集成测试 operators/test/integration/test_gemm_epilogue_fusion.py 展示了最简融合用法——把一元激活(relu/tanh/sigmoid/exp 等)以字符串形式的epi函数传入:
def epi(accum): D = unary_op(accum) return D epi_str = f"def epi(accum): D = {unary_str}(accum); return D" epi_args = ops.EpilogueArguments(epi_str, D=D) args = ops.GemmArguments(A=A, B=B, out=D, accumulator_type=accumulator_type, epilogue=epi_args)随后ops.get_operators(args, target_sm=...)找到支持算子并run,以epi(A @ B)为参考校验正确性。测试还覆盖了二元运算(add/sub/mul)与一元/二元复合融合。
六、Operands:操作数抽象
6.1 Operand 基类
Operand(base.py)是所有操作数的抽象基类:
- 最简情形下封装单个张量;
- 复杂情形下封装多个张量共同表达一个逻辑操作数——例如
ScaledOperand用"量化张量 + 缩放张量"重建操作数的逻辑值。
它提供copy()(浅拷贝,不拷贝底层张量,@final禁止覆写)与抽象的__copy__,以及_convert_to_internal_types钩子。
6.2 DenseTensor
DenseTensor(operators/cutlass/operators/arguments/operand.py)封装一个简单的稠密张量,是唯一的字段tensor: TensorLike。它的__getattr__会代理到底层张量,因此DenseTensor实例可以像其包装的张量一样被读取.shape、.dtype等属性。
6.3 ScaledOperand:块缩放(Block-Scaled)操作数
ScaledOperand(operand.py)的逻辑值为scale * quantized:scale张量的每个元素按mode给定的块形状广播并乘到quantized的一个连续块上。它主要用于表达窄精度格式(OCP MXFP8/MXFP4、NVIDIA NVFP4)——数据以窄精度量化张量存储,配合独立缩放张量恢复动态范围。
字段一览:
| 字段 | 类型 | 含义 |
|---|---|---|
quantized | DenseTensor | 窄精度量化值张量,逻辑值 =scale * quantized |
scale | DenseTensor | 缩放因子张量,必须连续、元素数恰好等于numel_scale(...)、且已按swizzle命名布局排布;形状本身不做校验 |
mode | ScaleMode \| tuple[int, ...] | 每个缩放因子广播覆盖的块形状,通常由数据格式决定(MXFP8/MXFP4/NVFP4 有指定 mode) |
swizzle | ScaleSwizzleMode | 缩放张量的内存布局,通常由硬件架构决定(如 Blackwell 块缩放 MMA 要求特定 swizzle) |
mode既可以是ScaleMode枚举,也可以是裸(L, M, K)元组——例如(1, 1, 32)表示每个缩放因子沿 K 轴覆盖 32 个元素、沿 L 与 M 轴覆盖 1 个元素。
numel_scale 静态方法(operand.py)返回scale张量应具有的元素数。设V = ScaleMode.numel(mode)为每个缩放因子覆盖的量化元素数,quantized_shape = (L, outer, K)(outer 对 A 侧是 M、对 B 侧是 N),则:
SwizzleNone:L * outer * ceil_div(K, V)Swizzle32x4x4:L * round_up(outer, 128) * round_up(ceil_div(K, V), 4)
quantized_shape为 rank-2 时按L=1处理;非法 rank 或未识别的 swizzle 抛ValueError。
6.4 ScaleMode:缩放粒度枚举
ScaleMode(operand.py)枚举常用块缩放模式,每个成员的值为(batch, row, col)元组:
| 成员 | 块形状 | 典型用途 |
|---|---|---|
Blockwise1x16 | (1, 1, 16) | NVIDIA NVFP4(FP8 E4M3 缩放 dtype) |
Blockwise1x32 | (1, 1, 32) | OCP MXFP8 / MXFP4(E8M0 缩放 dtype) |
枚举还提供compare()静态方法(支持枚举与裸元组混合比较,且容忍不同长度——只要较长元组多余的前导位置全是 1,即(1, 1, 16) == (1, 16))与numel()(返回块体积,如numel(Blockwise1x32) == 32)。
6.5 ScaleSwizzleMode:缩放张量内存布局枚举
ScaleSwizzleMode(operand.py)声明缩放因子已按特定硬件布局存放:
SwizzleNone:ScaleMode隐含的自然顺序,即(L, M, K)操作数在(1, 1, V)模式下有L * M * (K // V)个缩放因子,每个mode块一个值;Swizzle32x4x4:Blackwelltcgen05.mma块缩放 MMA 要求的 1D 块缩放布局(MXFP8/MXFP4/NVFP4 GEMM 使用)。每个 tile 含 128x4 个缩放因子,128 行按 32 一组交织,与 Blackwell 张量核的 warp-group 结构匹配。
需要强调的是:Operator API只校验scale 张量连续且元素数符合(mode, swizzle)要求,不检查也不重排其值——写入该布局是量化器(producer)的职责。
6.6 ScaledOperand 使用示例与测试
operand.py 中的官方示例:
quantized_A = torch.randn(M, K, dtype=torch.float8_e4m3fn, device="cuda") scale_A = torch.randn(M, K // 32, dtype=torch.float8_e8m0fnu, device="cuda") A = ScaledOperand( quantized_A, scale_A, ScaleMode.Blockwise1x32, ScaleSwizzleMode.Swizzle32x4x4, )集成测试 operators/test/integration/test_blockscaled_gemm.py 演示了端到端流程:用ops.ScaledOperand.numel_scale((L, M, K), scale, swizzle)分配 scale 张量,构造GemmArguments后经ops.get_operators(fake_args, target_sm=...)发现算子、operator.compile(fake_args)编译,再以真实张量operator.run(args)执行(支持 FakeTensor 模式下先编译、再复用的两段式流程)。
单元测试 operators/test/unit/test_arguments.py 则系统性验证了numel_scale的各类边界:K 非整除(ceil_div)、M 与 K 块的对齐/补齐(round_up)、裸元组 mode 与枚举等价、rank-2 等价于L=1、非法 rank 抛ValueError等。
七、类型标记(Type Markers):TensorLike 与 NumericLike
cutlass.operators.typing模块(operators/cutlass/operators/typing.py)为参数/操作数字段注解提供两个类型标记:
7.1 TensorLike
TensorLike: TypeAlias = _SupportsDLPack | cute.Tensor | TensorWrapper含义(typing.py):
- 任何支持 DLPack 协议(
__dlpack__/__dlpack_device__)的张量:torch.Tensor、jax.Array、numpy.ndarray等; cutlass.cute.Tensor(CuTe DSL 宿主张量,不实现 DLPack,但由TensorWrapper原生处理);TensorWrapper本身。
其中_SupportsDLPack是@runtime_checkable的 Protocol(typing.py),可用isinstance做运行时检查。这一设计正是 Operator API"原生支持 PyTorch 及其他 DLPack 张量"的协议基础。
7.2 NumericLike
class NumericLike(Protocol): """Type marker for fields that accept numeric-like types. ..."""被NumericLike注解的字段接受cutlass.Numeric与torch.dtype(typing.py),在转换阶段被to_cutlass_type统一为cutlass.Numeric。
八、端到端实战示例
综合以上各节,一个完整的最小 GEMM 调用链(与 operators/README.md 及 media/docs/operators/overview.rst 的示例一致):
import cutlass.operators as ops import torch A, B, out = (torch.randn(128, 128, device="cuda", dtype=torch.float16) for _ in range(3)) # 1. 用 RuntimeArguments 子类表达"要做什么"与"操作数" args = ops.GemmArguments(A, B, out, accumulator_type=torch.float32) # 2. 算子发现:找到支持该参数、且能在目标 SM 上运行的 Operator operators = ops.get_operators(args, target_sm="100") # 3. JIT 编译并执行 operators[0].run(args)再叠加块缩放操作数与自定义 epilogue(参考 test_blockscaled_gemm.py 与 test_gemm_epilogue_fusion.py 的写法):
# 块缩放 GEMM args = ops.GemmArguments( A=ops.ScaledOperand(A_fp8, SFA, ops.ScaleMode.Blockwise1x32, ops.ScaleSwizzleMode.Swizzle32x4x4), B=ops.ScaledOperand(B_fp8, SFB, ops.ScaleMode.Blockwise1x32, ops.ScaleSwizzleMode.Swizzle32x4x4), out=D, accumulator_type=torch.float32, ) # 融合 ReLU epilogue(epilogue_fn 也可传入字符串源码) def epi(accum): D = torch.relu(accum) return D args = ops.GemmArguments( A=A, B=B, out=D, accumulator_type=torch.float32, epilogue=ops.EpilogueArguments(epi, D=D), )运行环境提示:nvidia-cutlass-operators处于 beta 阶段(接口可能变更),可通过pip install nvidia-cutlass-operators[torch]安装(见 operators/README.md);示例与更多教程可参阅 operators/examples 下的 notebook。
九、小结
本文围绕 arguments.rst 展开,梳理了 CUTLASS Operator API 的完整参数体系:
- RuntimeArguments作为抽象基类定义"操作类型 + 操作数 + 性能控制"的骨架,
__post_init__驱动TensorLike → TensorWrapper、NumericLike → cutlass.Numeric的内部类型转换; - GemmArguments / GroupedGemmArguments分别承载稠密 GEMM(含便捷构造、问题尺寸派生与形状校验)与分组 GEMM(含 offsets 机制);
- EpilogueArguments通过 AST 解析把普通 Python 函数降为 EFC kernel 的 epilogue,并支持
Load/Store/Transport逐张量控制数据搬运策略; - Operand 家族(
DenseTensor、ScaledOperand与ScaleMode/ScaleSwizzleMode)统一了稠密与块缩放(MXFP8/MXFP4/NVFP4)两类操作数的表达; - TensorLike / NumericLike类型标记让框架张量(PyTorch、DLPack 等)与 CUTLASS 数值类型无缝对接。
理解这套参数对象,是使用 Operator API 完成"发现算子 → 编译 → 运行"全流程的第一步,也是接入自定义 CuTe DSL kernel(bring-your-own-kernel)与编写可移植集成代码的基础。后续可结合 api_reference/index.rst 中并列的 operator、discovery、metadata 等参考页,以及 media/docs/operators/tutorials/index.rst 下的 step-by-step 教程,构建完整的能力地图。
【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考