CUTLASS Operator API 参数体系完全解析:RuntimeArguments、Operands 与类型标记

发布时间:2026/9/16 17:21:57
CUTLASS Operator API 参数体系完全解析:RuntimeArguments、Operands 与类型标记 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协议的任何对象TensorLiketorch.Tensor、cute.Tensor、cutlass.operators.utils.tensor.TensorWrapper等NumericLikecutlass.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 中的抽象 dataclassdataclass class RuntimeArguments: Describes the operands and all other arguments passed to an Operator at runtime. ... performance: PerformanceControls | None field(defaultNone, kw_onlyTrue) 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_onlyTrue关键字专用指向PerformanceControls实例触发内部类型转换__post_init__中调用_convert_to_internal_types(self)把用户传入的框架张量、torch.dtype等统一转换为内部表示。2.2 PerformanceControlsPerformanceControls同文件 base.py是所有运行时性能选项的通用容器本身不强制任何字段——不同的算子/实现可以为其定义具体的子类字段如 tile 调度、工作区等从而在不改变参数接口的前提下支持运行时性能调优。2.3 类型自动转换机制核心原理_convert_to_internal_typesbase.py是整套参数体系的转换引擎。它利用 dataclass 的类型注解get_type_hints逐字段检查字段注解转换目标说明TensorLikeTensorWrapper用TensorWrapper(value, **global_metadata)包装负责编译期/运行期张量描述见下文NumericLikecutlass.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保证逻辑布局一致。三、GemmArgumentsGEMM 操作参数3.1 字段与便捷构造GemmArguments定义在 operators/cutlass/operators/arguments/gemm.py表示out A B的 GEMM 运算字段类型含义AOperand输入张量 Ashape(L, M, K)或(M, K)BOperand输入张量 Bshape(L, K, N)或(K, N)outOperand输出张量 Cshape(L, M, N)或(M, N)accumulator_typeNumericLike累加器数据类型epilogueEpilogueArguments \| None可选的 GEMM 后自定义 epilogue 融合performancePerformanceControls \| None继承自基类的运行时性能控制关键字参数其中M A 与 out 的行数N B 与 out 的列数K A 的列数与 B 的行数L 批量矩阵乘法批数。所有张量必须同为 rank-3 或同为 rank-2。便捷构造对稠密 GEMMA、B、out可以直接传裸张量而不必包DenseTensorGemmArguments(A, B, out, accumulator_type) # 等价于 GemmArguments(DenseTensor(A), DenseTensor(B), DenseTensor(out), accumulator_type)而其他操作数类型必须显式包装。例如块缩放 GEMMscaled GEMMGemmArguments( ScaledOperand(A, ScaleATensor, scale_mode, scale_swizzle), ScaledOperand(B, ScaleBTensor, scale_mode, scale_swizzle), out, # 无需 DenseTensor 包装 accumulator_type, )该便捷逻辑由_operand_or_denseoperand.py实现若传入对象已是Operand则直接返回若是TensorLike则包装为DenseTensor否则抛出ValueError。3.2 问题尺寸派生GemmProblemSizeGemmProblemSizegemm.py是一个NamedTuple其(M, N, K, L)与操作数形状对应关系为A(L, M, K)B(L, K, N)out(L, M, N)GemmArguments.problem_size属性直接由操作数形状推导gemm.pyM A.shape[-2]、N B.shape[-1]、K A.shape[-1]、L A.shape[0] if rank3 else 1。这意味着算子发现ops.get_operators无需用户显式声明问题尺寸全部从张量自描述。3.3 参数校验_validate构造时GemmArguments会调用_validate()gemm.py执行以下一致性检查不通过即抛出带完整形状信息的ValueErrorA、B、out 必须是 rank 2 或 3A 的 K 维与 B 的 K 维相等注意_logical_contraction_extent会考虑亚字节 dtype 的 K 维打包系数见 gemm.pyout 的 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变体Ashape(TotalM, K)或(1, TotalM, K)Bshape(problems_in_group, K, N)outshape(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基础上增加了offsetsOperand类型dataclass 字段带有metadata{alignment_bytes: 4}即要求 4 字节对齐供TensorWrapper使用。构造函数签名grouped_gemm.pyGroupedGemmArguments( A, B, out, accumulator_type, offsetsoffsets, epilogueNone, )集成测试 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 约束EpilogueArgumentsoperators/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_layoutepilogue.py以保证 EFC kernel 的跨 CTA atomic 归约路径可用。5.4 Load / Store / Transport逐张量数据搬运策略默认情况下epilogue 的张量通过TMA搬运。EpilogueArguments支持用Load/Store描述符包装 kwargs 来覆盖默认搬运方式epilogue.pyCops.Load(C, viaops.Transport.ASYNC_GMEM_LOAD) Dops.Store(D, viaops.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_LOADStore 仅允许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 fdef epi(accum): D {unary_str}(accum); return D epi_args ops.EpilogueArguments(epi_str, DD) args ops.GemmArguments(AA, BB, outD, accumulator_typeaccumulator_type, epilogueepi_args)随后ops.get_operators(args, target_sm...)找到支持算子并run以epi(A B)为参考校验正确性。测试还覆盖了二元运算add/sub/mul与一元/二元复合融合。六、Operands操作数抽象6.1 Operand 基类Operandbase.py是所有操作数的抽象基类最简情形下封装单个张量复杂情形下封装多个张量共同表达一个逻辑操作数——例如ScaledOperand用量化张量 缩放张量重建操作数的逻辑值。它提供copy()浅拷贝不拷贝底层张量final禁止覆写与抽象的__copy__以及_convert_to_internal_types钩子。6.2 DenseTensorDenseTensoroperators/cutlass/operators/arguments/operand.py封装一个简单的稠密张量是唯一的字段tensor: TensorLike。它的__getattr__会代理到底层张量因此DenseTensor实例可以像其包装的张量一样被读取.shape、.dtype等属性。6.3 ScaledOperand块缩放Block-Scaled操作数ScaledOperandoperand.py的逻辑值为scale * quantizedscale张量的每个元素按mode给定的块形状广播并乘到quantized的一个连续块上。它主要用于表达窄精度格式OCP MXFP8/MXFP4、NVIDIA NVFP4——数据以窄精度量化张量存储配合独立缩放张量恢复动态范围。字段一览字段类型含义quantizedDenseTensor窄精度量化值张量逻辑值 scale * quantizedscaleDenseTensor缩放因子张量必须连续、元素数恰好等于numel_scale(...)、且已按swizzle命名布局排布形状本身不做校验modeScaleMode \| tuple[int, ...]每个缩放因子广播覆盖的块形状通常由数据格式决定MXFP8/MXFP4/NVFP4 有指定 modeswizzleScaleSwizzleMode缩放张量的内存布局通常由硬件架构决定如 Blackwell 块缩放 MMA 要求特定 swizzlemode既可以是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则SwizzleNoneL * outer * ceil_div(K, V)Swizzle32x4x4L * round_up(outer, 128) * round_up(ceil_div(K, V), 4)quantized_shape为 rank-2 时按L1处理非法 rank 或未识别的 swizzle 抛ValueError。6.4 ScaleMode缩放粒度枚举ScaleModeoperand.py枚举常用块缩放模式每个成员的值为(batch, row, col)元组成员块形状典型用途Blockwise1x16(1, 1, 16)NVIDIA NVFP4FP8 E4M3 缩放 dtypeBlockwise1x32(1, 1, 32)OCP MXFP8 / MXFP4E8M0 缩放 dtype枚举还提供compare()静态方法支持枚举与裸元组混合比较且容忍不同长度——只要较长元组多余的前导位置全是 1即(1, 1, 16) (1, 16)与numel()返回块体积如numel(Blockwise1x32) 32。6.5 ScaleSwizzleMode缩放张量内存布局枚举ScaleSwizzleModeoperand.py声明缩放因子已按特定硬件布局存放SwizzleNoneScaleMode隐含的自然顺序即(L, M, K)操作数在(1, 1, V)模式下有L * M * (K // V)个缩放因子每个mode块一个值Swizzle32x4x4Blackwelltcgen05.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, dtypetorch.float8_e4m3fn, devicecuda) scale_A torch.randn(M, K // 32, dtypetorch.float8_e8m0fnu, devicecuda) 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 等价于L1、非法 rank 抛ValueError等。七、类型标记Type MarkersTensorLike 与 NumericLikecutlass.operators.typing模块operators/cutlass/operators/typing.py为参数/操作数字段注解提供两个类型标记7.1 TensorLikeTensorLike: TypeAlias _SupportsDLPack | cute.Tensor | TensorWrapper含义typing.py任何支持 DLPack 协议__dlpack__/__dlpack_device__的张量torch.Tensor、jax.Array、numpy.ndarray等cutlass.cute.TensorCuTe DSL 宿主张量不实现 DLPack但由TensorWrapper原生处理TensorWrapper本身。其中_SupportsDLPack是runtime_checkable的 Protocoltyping.py可用isinstance做运行时检查。这一设计正是 Operator API原生支持 PyTorch 及其他 DLPack 张量的协议基础。7.2 NumericLikeclass NumericLike(Protocol): Type marker for fields that accept numeric-like types. ...被NumericLike注解的字段接受cutlass.Numeric与torch.dtypetyping.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, devicecuda, dtypetorch.float16) for _ in range(3)) # 1. 用 RuntimeArguments 子类表达要做什么与操作数 args ops.GemmArguments(A, B, out, accumulator_typetorch.float32) # 2. 算子发现找到支持该参数、且能在目标 SM 上运行的 Operator operators ops.get_operators(args, target_sm100) # 3. JIT 编译并执行 operators[0].run(args)再叠加块缩放操作数与自定义 epilogue参考 test_blockscaled_gemm.py 与 test_gemm_epilogue_fusion.py 的写法# 块缩放 GEMM args ops.GemmArguments( Aops.ScaledOperand(A_fp8, SFA, ops.ScaleMode.Blockwise1x32, ops.ScaleSwizzleMode.Swizzle32x4x4), Bops.ScaledOperand(B_fp8, SFB, ops.ScaleMode.Blockwise1x32, ops.ScaleSwizzleMode.Swizzle32x4x4), outD, accumulator_typetorch.float32, ) # 融合 ReLU epilogueepilogue_fn 也可传入字符串源码 def epi(accum): D torch.relu(accum) return D args ops.GemmArguments( AA, BB, outD, accumulator_typetorch.float32, epilogueops.EpilogueArguments(epi, DD), )运行环境提示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 kernelbring-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),仅供参考