CUTLASS Python 接口示例全解析:从 Basic GEMM、Epilogue 到 PyTorch CUDA 扩展
2026/9/16 16:20:05 网站建设 项目流程

CUTLASS Python 接口示例全解析:从 Basic GEMM、Epilogue 到 PyTorch CUDA 扩展

【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass

本文是 CUTLASS Python 接口「示例与教程(Examples and Tutorials)」部分的完整技术指南。它以 python/docs_src/source/examples.rst 为骨架,逐篇拆解该文档编排的三个 Jupyter 示例:Basic GEMM(基础矩阵乘)Epilogue(尾处理与逐元素激活函数)Grouped GEMM 导出为 PyTorch CUDA 扩展。读完本文,你将掌握如何用 Python 声明、编译、运行 CUTLASS GEMM 内核,如何切换 Tensor Core / SIMT 计算模式、融合 ReLU 等激活函数,以及如何把内核导出为可 JIT 编译的 PyTorch 扩展并做性能对比。

背景:三个示例在文档体系中的位置

examples.rst是 CUTLASS Python 接口 Sphinx 文档中「Examples and Tutorials」章节的入口,其核心内容是一张toctree表,将读者导向三份 Notebook 文件:

.. toctree:: :maxdepth: 5 Basic GEMM <externals/00_basic_gemm.nblink> Epilogue <externals/01_epilogue.nblink> PyTorch Extension <externals/02_pytorch_extension_grouped_gemm.nblink>

三份.nblink是 Sphinx 的 Notebook 链接文件,分别指向对应的.ipynb(见 00_basic_gemm.nblink、01_epilogue.nblink、02_pytorch_extension_grouped_gemm.nblink):

  • 00_basic_gemm.ipynb:声明、编译并运行基础 GEMM;
  • 01_epilogue.ipynb:为 GEMM 融合各种逐元素激活函数;
  • 02_pytorch_extension_grouped_gemm.ipynb:把 Grouped GEMM 导出为 PyTorch CUDA 扩展。

仓库中这些 Notebook 的原始源码位于 examples/python/deprecated/ 目录(00_basic_gemm.ipynb01_epilogue.ipynb02_pytorch_extension_grouped_gemm.ipynb),docs 目录下保存的是构建产物。三份 Notebook 展示的核心 API 形态为cutlass.op.Gemmcutlass.op.GroupedGemmcutlass.epilogue.<act>cutlass.emit.pytorch,与当前仓库中 python/cutlass_cppgen/op/gemm.py、python/cutlass_cppgen/op/gemm_grouped.py、python/cutlass_cppgen/emit/pytorch.py 的实现一一对应,下文将结合这些源码逐层剖析。

运行环境与安装准备

三个示例都运行在安装了 CUTLASS Python 接口的环境中。根据 python/docs_src/source/install.md 的说明,安装方式有两种:

  • 安装稳定版:pip install nvidia-cutlass(注意:PyPI 上其他名为cutlass的包与 NVIDIA CUTLASS 无关);
  • 从源码安装:在 CUTLASS 仓库根目录执行pip install .(开发者模式用pip install -e .),要求本机安装的 CUDA Toolkit 与cuda-python的 major.minor 版本匹配。

安装前可选的几个环境变量及其推断规则如下:

环境变量作用未设置时的推断规则
CUTLASS_PATHCUTLASS 仓库路径当前目录上一级(本地安装),或cutlass_library安装位置的source目录
CUDA_INSTALL_PATHCUDA 安装路径第一个nvcc所在的/bin/nvcc上级目录(即which nvcc的结果)

也可以直接使用 NGC PyTorch Docker 容器快速上手:docker run --gpus all -it --rm nvcr.io/nvidia/pytorch:23.08-py3

示例一:Basic GEMM —— 声明、编译、运行一次搞定

00_basic_gemm.ipynb演示了 CUTLASS Python 接口「以最少配置跑通 GEMM」的核心工作流。

构造输入张量

首先导入依赖并构造 fp16 的输入/输出张量:

import numpy as np import random import cutlass # 控制是否在每一步打印生成的 C++ GEMM 声明,设为 False 可省略输出 print_module = True m = 128 n = m k = m dtype = np.float16 type_A = np.float16 type_B = np.float16 type_C = np.float16 type_D = np.float16 np.random.seed(1234) random.seed(1234) scope_min = -4 scope_max = 4 tensor_A = np.ceil(np.random.uniform(low=scope_min, high=scope_max, size=(m, k)).astype(type_A)) tensor_B = np.ceil(np.random.uniform(low=scope_min, high=scope_max, size=(k, n)).astype(type_B)) tensor_C = np.ceil(np.random.uniform(low=scope_min, high=scope_max, size=(m, n)).astype(type_C)) alpha = np.float16(1.) beta = np.float16(0.) tensor_D = np.zeros(tensor_C.shape).astype(type_D)

这里使用固定随机种子(np.random.seed(1234))保证示例可复现;np.ceil把随机浮点取整,便于后续与 NumPy 结果做精确相等比较。

声明并运行默认 GEMM

只需把张量交给cutlass.Gemm,接口就会为当前设备挑选一套默认的 GEMM 配置:

# 显式指定 element_accumulator,使其与后面 NumPy 参考实现的累加类型一致; # 若累加类型与 element 相同,则可以不指定 plan = cutlass.Gemm(element=dtype, layout=cutlass.LayoutType.RowMajor, element_accumulator=np.float32) plan.run(tensor_A, tensor_B, tensor_C, tensor_D, print_module=print_module)

plan.run()的调用链路是「生成 CUTLASS C++ 内核 → 编译 → 在传入张量上执行」。print_module=True时会在屏幕上打印生成的 C++ 代码。从 Notebook 的输出来看,默认(假设运行在 SM80)会生成一个基于 FP16 Tensor Core 的 kernel,例如:

// Gemm operator cutlass_sm80_tensorop_f16_s16x8x16gemm_f16_1x1x1_256x128_64x3_tt_align8 using cutlass_sm80_tensorop_f16_s16x8x16gemm_f16_1x1x1_256x128_64x3_tt_align8_base = typename cutlass::gemm::kernel::DefaultGemmUniversal<...>

从这个 kernel 名称可以读出完整配置:SM80 架构、TensorOp、fp16、指令形状16x8x16、线程块形状256x128、K 步长64、3 个流水级、tt(A/B 均为 row-major)、对齐 8。这些默认参数正是由 python/cutlass_cppgen/op/gemm.py 中Gemm.__init__construct()的自动配置逻辑挑选出来的——construct()会根据数据类型推导 A/B 的最优对齐(min(128 // DataTypeSize[...], max(alignments("A")))),并在未指定tile_description时从可能的操作集中选取第一个配置(见 python/cutlass_cppgen/op/gemm.py#L417-L477)。

用 NumPy 校验结果

示例用 NumPy 逐元素比对,验证 kernel 正确性:

tensor_D_numpy = (alpha * (tensor_A @ tensor_B)) + (beta * tensor_C) np.testing.assert_array_equal(tensor_D, tensor_D_numpy)

值得注意的是,同一个 kernel 声明可以复用于其他框架(PyTorch、CuPy 等)提供的张量——接口在运行时通过_verify_tensor校验传入张量的数据类型与布局(见Gemm.run的实现),只要类型布局一致即可直接复用。

切换计算模式:TensorOp 与 Simt

默认情况下接口优先使用 Tensor Core(TensorOp);若配置在 Tensor Core 上不受支持,则自动回退到 SIMT kernel。当前使用的操作模式可通过plan.opclass属性查询:

print(plan.opclass) # Tensor Core 操作

如果想强制使用 CUTLASS 的 SIMT GEMM,只需改写opclass字段:

tensor_D_simt = np.zeros(tensor_C.shape).astype(type_D) plan.opclass = cutlass.OpcodeClass.Simt plan.run(tensor_A, tensor_B, tensor_C, tensor_D_simt, alpha, beta, print_module=print_module)

此时打印出的 kernel 模板参数会切换为 CUTLASS SIMT GEMM 的形式。再次用np.testing.assert_array_equal(tensor_D, tensor_D_simt)可确认 Tensor Core 与 SIMT 两种实现的结果完全一致。这一机制在源码中的对应点是OperationBase.opclass的 getter/setter(python/cutlass_cppgen/op/op.py#L209-L221),以及Gemm.__init__中「能支持 TensorOp 就用 TensorOp,否则回退 Simt」的默认逻辑(python/cutlass_cppgen/op/gemm.py#L289-L296)。

内核缓存:避免重复编译

示例特意提醒:前两次plan.run()耗时较长,是因为内核尚未编译。CUTLASS 会缓存已编译的二进制,同一内核再次运行(哪怕换了更大的张量、不同的 alpha/beta)无需重新编译。比如把问题规模放大到2400 x 3232 x 4096并把beta改为2.后再次plan.run(...),编译开销不再出现:

m = 2400 n = 3232 k = 4096 # ... 重新生成 tensor_A/B/C/D,alpha = 1., beta = 2. plan.opclass = cutlass.OpcodeClass.TensorOp plan.run(tensor_A, tensor_B, tensor_C, tensor_D, alpha, beta, print_module=print_module)

运行非默认配置:tile_descriptions 与 compile

默认配置只是「开箱即用」的选择。需要更精细控制时,plan.tile_descriptions()会返回接口从 CUTLASS profiler 枚举出的全部合法配置(即所有可行的 tile 形状、指令形状与流水级组合):

tiles = plan.tile_descriptions() print('{} tile descriptions returned'.format(len(tiles))) num_print = 10 print('First {} tile descriptions are:'.format(num_print)) for td in tiles[:num_print]: print(td)

然后可以任选其中一个配置单独编译、运行:

idx = random.randint(0, len(tiles)-1) td = tiles[idx] print('Tile description {} is: {}'.format(idx, td)) plan.compile(td) plan.run(tensor_A, tensor_B, tensor_C, tensor_D, alpha, beta, print_module=print_module)

其中plan.compile(td)compilerun解耦,便于「先编译、后执行」的工作流。在源码中,tile_descriptions()Gemm类实现(python/cutlass_cppgen/op/gemm.py#L405-L410),它从possible_operations.all_operations中把每个 profiler 操作转成TileDescription;而compile()则负责把 tile 描述与张量对齐等信息绑定成GemmOperationUniversal并调用后端编译器生成、编译内核。

更换 Swizzling:Stream K 示例

接口还允许修改内核的 swizzling 函数。例如切换到 CUTLASS 的Stream K特性:

# Stream K 仅在 SM90 之前受支持(至少在本示例写作时如此) if plan.cc != 90: plan.swizzling_functor = cutlass.swizzle.ThreadblockSwizzleStreamK plan.run(tensor_A, tensor_B, tensor_C, tensor_D, alpha, beta, print_module=print_module)

源码层面的约束与此一致:swizzling_functor的 setter 会校验ThreadblockSwizzleStreamK只能配合 TensorOp 使用、且在 SM90+ 上不受支持(python/cutlass_cppgen/op/gemm.py#L319-L328);ThreadblockSwizzleStreamK定义在 python/cutlass_cppgen/swizzle.py。

错误处理:共享内存不足的友好报错

CUTLASS Python 接口会尽量把运行期/编译期错误在 Python 层捕获,给出更易读的报错信息。示例演示的场景是:为一个 GEMM 设置过多流水级(stages),导致 GPU 共享内存不足以启动 kernel。若直接调用 C 层会得到晦涩的运行时错误,而接口会拦截并提示:

# td = tiles[0] # td.stages = 8 # plan.compile(td)

run()中通过super().run_setup()等前置检查与_valid_tile_description校验(见construct()中对非法 tile 描述抛出"Invalid tile description."的分支),把配置合法性检查前置到 Python 层,正是「更可理解的错误信息」的来源。

示例二:Epilogue —— 一行代码融合激活函数

01_epilogue.ipynb演示如何为 GEMM 融合逐元素激活函数。示例仍以m = n = k = 256的 fp16 GEMM 为背景构造张量(过程与示例一相同,此处不再重复)。

默认的 identity 尾处理

直接运行默认 GEMM,执行的就是标准的线性组合:

plan = cutlass.op.Gemm(element=np.float16, layout=cutlass.LayoutType.RowMajor) plan.run(tensor_A, tensor_B, tensor_C, tensor_D, print_module=print_module)

默认激活函数是identity(恒等),无需显式指定,对应的数学形式为:

D = alpha * (A @ B) + beta * C

融合 ReLU

在 GEMM 的线性组合之后再做一次逐元素变换。设激活函数为act,最终形式为:

D = alpha * (A @ B) + beta * C D = act(D)

CUTLASS 中融合 ReLU 只需设置 plan 的activation字段:

tensor_D_relu = np.zeros(tensor_C.shape).astype(type_D) plan.activation = cutlass.epilogue.relu plan.run(tensor_A, tensor_B, tensor_C, tensor_D_relu, print_module=print_module)

其中 ReLU 对输入x返回max(x, 0)。随后用 NumPy 精确校验:

relu_ref = (tensor_D >= 0).astype(type_D) * tensor_D np.testing.assert_array_equal(relu_ref, tensor_D_relu)

更多激活函数:plan.activations()

接口内置了一批常用逐元素激活函数,可通过plan.activations()列出并逐个运行:

activations = plan.activations() for activation in activations: print(activation) for activation in activations: print('=============================================================================================') print(f'Compiling and running activation {activation}') print('=============================================================================================') plan.activation = activation plan.run(tensor_A, tensor_B, tensor_C, tensor_D, print_module=print_module)

结合源码可以确认这些激活函数的完整清单。OperationBase.activations()直接返回get_activations()(python/cutlass_cppgen/op/op.py#L106-L110),而注册表定义在 python/cutlass_cppgen/epilogue/epilogue.py:

_activations = [gelu, hardswish, identity, leaky_relu, relu, sigmoid, silu, tanh]

即内置激活函数为:gelu、hardswish、identity、leaky_relu、relu、sigmoid、silu、tanh,全部通过cutlass.epilogue.<name>导入(python/cutlass_cppgen/epilogue/init.py)。这些激活函数最终会被翻译为 C++ epilogue functor(get_activation_epilogue负责按输出数据类型、对齐等参数构造对应的 epilogue 实现),从而做到「零额外 kernel 启动开销」的算子融合。

示例三:Grouped GEMM 导出为 PyTorch CUDA 扩展

02_pytorch_extension_grouped_gemm.ipynb展示了从「快速实验」到「生产接入 PyTorch」的完整路径。

Grouped GEMM 是什么

Grouped GEMM 允许在单个 CUDA kernel内执行一组 GEMM,其中每个 GEMM 可以有不同的尺寸和 stride。它可以看作指针数组 GEMM 的泛化形式——不要求各 GEMM 的尺寸与 stride 相同。例如有p个 GEMM,尺寸分别为:

M_1 x N_1 x K_1 M_2 x N_2 x K_2 ... M_p x N_p x K_p

它们可以在一次 kernel launch 中被统一调度执行,避免为每个小 GEMM 单独启动 kernel 带来的开销。

声明 GroupedGemm 并批量运行

示例使用 PyTorch 张量构造 fp16 的 Grouped GEMM:

import cutlass import torch dtype = torch.float16 plan = cutlass.op.GroupedGemm(element=dtype, layout=cutlass.LayoutType.RowMajor)

随后是两组工具函数:initialize(dtype, M, N, K)为单个 GEMM 生成 A、B、C、D 四个张量;generate_problems(problems)[128, 256, 512, 1024]中随机挑选尺寸,生成一批 GEMM:

import random random.seed(2023) # Utility function to initialize A, B, C, and D matrices corresponding to dimensions M, N, and K def initialize(dtype, M, N, K): sizes = [(M, K), (K, N), (M, N), (M, N)] return [torch.randint(-3, 3, size, device='cuda').to(dtype) for size in sizes] # Utility function to generate `problems` GEMMs of random sizes def generate_problems(problems): valid_sizes = [128, 256, 512, 1024] As, Bs, Cs, Ds = [], [], [], [] for _ in range(problems): M, N, K = [random.choice(valid_sizes) for _ in range(3)] A, B, C, D = initialize(dtype, M, N, K) As.append(A); Bs.append(B); Cs.append(C); Ds.append(D) return As, Bs, Cs, Ds

对一组 50 个 GEMM 批量运行,并与 PyTorch 的逐对matmul结果对比:

As, Bs, Cs, Ds, = generate_problems(50) plan.run(As, Bs, Cs, Ds, print_module=True) Ds_torch = [a @ b for a, b in zip(As, Bs)] for d, d_torch in zip(Ds, Ds_torch): assert torch.allclose(d, d_torch)

在源码中,GroupedGemm继承自Gemm(python/cutlass_cppgen/op/gemm_grouped.py#L73),并针对 SM90+ 做了一次降级:由于 Grouped GEMM 的 SM90 特化当前不可用,构造时会把配置回退到 SM80(if self.current_cc in [90, 100, 101, 103]: self._reset_options(80))。此外它不支持更换 swizzling functor(setter 直接抛异常),这些都是使用该操作时需要注意的边界。

导出为 PyTorch CUDA 扩展

Python 直跑适合快速实验,但生产接入时更倾向通过 PyTorch CUDA 扩展使用 CUTLASS kernel,以去掉 Python 层带来的运行时开销。接口提供了两种生成方式:写盘供「预编译(ahead-of-time)」,或直接 JIT 编译返回给用户。JIT 方式只需三步:

op = plan.construct() grouped_gemm = cutlass.emit.pytorch(op, name='grouped_gemm', cc=plan.cc, sourcedir='out', jit=True)

cutlass.emit.pytorch会在out/目录下生成三个文件:

  • out/grouped_gemm_kernel.cu:CUTLASS kernel 的声明,以及从 PyTorch 张量调用它的方法;
  • out/grouped_gemm.cpp:对上述 CUTLASS kernel 的 C++ 封装;
  • setup.py:用于构建并安装该扩展的setuptools脚本。

jit=True时扩展会被即时编译、加载并直接返回;若jit=False,则只把源码写入sourcedir供后续手动构建。对应实现见 python/cutlass_cppgen/emit/pytorch.py#L905-L928:函数会根据操作类型(GemmOperationUniversal/GemmOperationGrouped/Conv2dOperation)分派到不同的发射器,分别产出<name>_kernel.cu(含setup.pyextra_compile_args)等文件。

out/目录下手动构建扩展的方式(AOT 场景)为:

TORCH_CUDA_ARCH_LIST="8.0" python setup.py install

其中TORCH_CUDA_ARCH_LIST需设置为运行该 kernel 的设备的 compute capability(例如 H100 用9.0,A100 用8.0)。

运行扩展并做性能对比

加载后的扩展用法与普通 PyTorch 模块一致:

Ds = grouped_gemm.run(As, Bs) Ds_torch = [a @ b for a, b in zip(As, Bs)] for d, d_torch in zip(Ds, Ds_torch): assert torch.allclose(d, d_torch)

最后是标准的「预热 + 计时」性能对比流程:20 次 warmup,100 次计时,分别统计 Grouped GEMM 扩展与逐对 PyTorchmatmul的耗时,并打印二者比值:

num_warmup = 20 num_profile = 100 # Warmup iterations for _ in range(num_warmup): Ds = grouped_gemm.run(As, Bs) Ds_torch = [a @ b for a, b in zip(As, Bs)] torch.cuda.synchronize() # Timing iterations import time grouped = 0 nongrouped = 0 for _ in range(num_profile): start = time.time() Ds = grouped_gemm.run(As, Bs) torch.cuda.synchronize() grouped += time.time() - start start = time.time() Ds_torch = [a @ b for a, b in zip(As, Bs)] torch.cuda.synchronize() nongrouped += time.time() - start print('Grouped: {:.3f} us'.format(grouped * 1e6 / num_profile)) print('Non-Grouped: {:.3f} us'.format(nongrouped * 1e6 / num_profile)) print('Speedup: {:.3f}x'.format(nongrouped / grouped))

需要强调的是,该对比是「单 kernel 批量调度」对「逐个 kernel 启动」的工程性比较,具体加速比取决于问题规模分布与硬件环境,示例本身不承诺固定数值。

机制小结:示例背后的统一调用链

三个示例尽管场景不同,背后共享同一条调用链,可以在源码中完整追踪:

  1. 构造Gemm/GroupedGemm构造器绑定 A/B/C/D 的数据类型与布局(支持element+layout简写、逐操作数element_A/layout_A细粒度指定、或直接传代表性张量三种方式,优先级为「张量 > 逐操作数参数 > 通用参数」,见 python/cutlass_cppgen/op/gemm.py#L140-L260);
  2. construct():根据对齐偏好与默认 tile 生成GemmOperationUniversal/GemmOperationGrouped操作对象;
  3. compile():把操作对象交给后端编译器生成 C++ 源码并编译,产物按内核签名缓存,避免重复编译;
  4. run():校验运行时张量 → 复用或触发编译 → 计算 problem size / batch → 在指定 CUDA stream 上 launch,默认同步等待完成(sync=False时返回的GemmArguments可稍后手动sync())。

这套设计让「一行代码跑 GEMM」「切 opclass / swizzling / activation 反复试验」「把调好的内核一键导出为 PyTorch 扩展」成为可能。更底层的 C++ kernel 生成逻辑(DefaultGemmUniversal等模板声明)则落在 include/cutlass/gemm/kernel/ 与python/cutlass_library/(如 gemm_operation.py 中的GroupedGemmOperation)中,感兴趣的读者可以顺着这条链继续深入。

【免费下载链接】cutlassCUDA Templates and Python DSLs for High-Performance Linear Algebra项目地址: https://gitcode.com/GitHub_Trending/cu/cutlass

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询