拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyTorch 中的 CUTLASS 扩展库:从 FasterTransformer 移植的 fp16/bf16 × int8/int4 混合精度 GEMM 支持

PyTorch 中的 CUTLASS 扩展库从 FasterTransformer 移植的 fp16/bf16 × int8/int4 混合精度 GEMM 支持【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch本文围绕 PyTorch 源码树中的 cutlass_extensions 目录 展开说明这份从 NVIDIA FasterTransformer 项目移植的 CUTLASS 扩展代码的来龙去脉、目录结构、为适配 CUTLASS 3.x 所做的关键改动以及它如何支撑 MixedDtypesLinear.cu 中注册的_mixed_dtypes_linear算子完成“浮点激活 × 整型权重量化”的线性层计算。读完本篇你将理解这份扩展库在 PyTorch 混合精度推理栈中的定位、其头文件的分层职责以及调用链从 Python 算子入口到 CUTLASS kernel 的完整路径。目录定位为混合数据类型 GEMM 服务的移植代码README 开宗明义地说明aten/src/ATen/native/cuda/cutlass_extensions目录中的文件复制自 FasterTransformer 项目的src/fastertransformer/cutlass_extensions/include/cutlass_extensions目录其唯一目的就是支撑 PyTorch 中 MixedDTypesLinear.cu 文件的混合数据类型mixed datatypesGEMM 实现。文档还强调了三点关键事实只复制了必要文件并非 FasterTransformer 中该目录的全部内容都搬了过来仅保留了 PyTorch 该功能所需的子集改动最小化目标明确原始拷贝取自 FasterTransformer 的f8e42aa提交而该项目当时基于 CUTLASS 2.10因此 PyTorch 侧的改动核心是适配CUTLASS 3.xLint 导致外观差异拷贝到 PyTorch 后的文件按照 PyTorch 的 lint 规则重新格式化因此与原始文件在外观上差异较大。为了追踪真实改动README 直接内附了一份 lint 之前两套文件的 diff见下文逐条解读。这份 README 的写作方式很典型它不是一个功能说明书而是一份“移植档案”把上游出处、对应提交、改动边界、以及“为什么要保留这份外部代码”都记录了下来。目录结构与各头文件职责结合目录实际内容与 PyTorch 的调用点这套扩展头文件按 CUTLASS 的抽象层级组织aten/src/ATen/native/cuda/cutlass_extensions/ ├── README.md ├── arch/mma.h # 架构层带 dequantize 的 MMA 操作封装 ├── epilogue/thread/ft_fused_activations.h # 尾声层FasterTransformer 风格的融合激活 ├── epilogue_helpers.h # Epilogue 标签分发Bias/ReLU/SiLU 等 ├── ft_gemm_configs.h # GEMM tile/split-K 配置枚举 ├── interleaved_numeric_conversion.h # 交错的数值类型转换如 uint4 解包 ├── tile_interleaved_layout.h # ColumnMajorTileInterleave 布局定义 └── gemm/ ├── kernel/ │ ├── fpA_intB_gemm.h # 核心 kernel 模板 GemmFpAIntB │ ├── default_fpA_intB_traits.h # fpA_intB GEMM 的默认 traits │ └── mixed_gemm_B_layout.h # 混合 GEMM 的 B 矩阵布局适配 ├── threadblock/ │ ├── default_mma.h / default_mma_bf16.h # threadblock MMA 默认配置含 bf16 变体 │ ├── default_dq_mma.h / _pipelined / _multistage # dequantize 版 MMA 配置 │ ├── dq_mma_multistage.h / dq_mma_pipelined.h / dq_mma_base.h └── warp/ ├── default_mma_tensor_op.h # warp 级 TensorOp MMA 默认配置 ├── mma_tensorop_compute_B_with_f16.h └── mma_tensorop_dequantizer.h # B 矩阵整型在线反量化几个值得注意的核心构件GemmFpAIntBkernelfpA_intB_gemm.h这是整个扩展库的主角A 矩阵为浮点fpB 矩阵为整型int的 GEMM kernel 模板。其Arguments结构除了标准的problem_size、ref_A、ref_B外还额外携带ref_scale逐行缩放因子引用、gather/scatter 索引等字段这正是“权重按行量化、在线反量化”语义的体现ColumnMajorTileInterleave布局tile_interleaved_layout.h一个仅含RowsPerTile与ColumnsInterleaved两个模板参数的布局标签类型配套IsColumnMajorTileInterleave类型特征trait。它描述了整型权重在内存中的交错interleaved存放方式以匹配 dequantizer 的批量读取模式Epilogue 标签分发epilogue_helpers.hfastertransformer命名空间下用空结构体EpilogueOpNoBias、EpilogueOpBias、EpilogueOpBiasReLU、EpilogueOpBiasSilu、EpilogueOpBiasFtGelu作为编译期标签配合Epilogue模板特化将标签映射到具体的 CUTLASS 线程级 epilogue 算子如LinearCombination、LinearCombinationRelu、LinearCombinationSilu且统一使用NoBetaScalingGEMM 配置枚举ft_gemm_configs.h保留了 FasterTransformer 的CutlassTileConfig如CtaShape32x128x64_WarpShape32x32x64等、SplitKStyle与CutlassGemmConfig定义。其中注释明确提醒做权重-only 量化时运行时配置的 K 形状必须与 kernel 布局细节中的 K 形状一致——而 PyTorch 的调用端恰好选用了32x128x64的 CTA 形状与这份枚举中的CtaShape32x128x64_WarpShape32x32x64相吻合。适配 CUTLASS 3.x 的关键改动README 所附 diff 解读README 中最有信息量的部分是那份 lint 前的原始 diff它精确刻画了 PyTorch 相对 FasterTransformer 原版的真实改动。主要有四处去掉 workspace 与信号量依赖fpA_intB_gemm.h 的Params原版的Params构造需要传入grid_tiled_shape、gemm_k_size和void* workspace并持有semaphore成员PyTorch 版改为传入device_sms与sm_occupancy并在构造时内部通过ThreadblockSwizzle计算grid_tiled_shape与gemm_k_size。同时新增三个方法get_workspace_size()恒返回 0、init_workspace()恒返回成功、get_grid_dims()由 swizzle 推导。可以推断这一改动的意义是PyTorch 使用串行 split-K 因子固定为 1 的配置无跨 CTA 归约需求从而省去了 FasterTransformer 用于 split-K 并行归约的 workspace 与 semaphore 基础设施也删除了get_extra_workspace_size静态方法。这与 MixedDtypesLinear.cu 中SplitKFactor 1并附注“!1 会输出错误”的约束互相印证新增invoke静态入口为GemmFpAIntB补充了CUTLASS_DEVICE static void invoke(Params const, SharedStorage)内部构造GemmFpAIntB op; op(params, shared_storage);调用。从源码结构看这是为对齐 CUTLASS 3.x 中DeviceAdapter::launch_kernel要求的 kernel 入口约定kernel 需可通过invoke静态函数统一调用BF16 支持的宏条件放宽mma_tensorop_dequantizer.h原版依赖 FasterTransformer 内部头cuda_bf16_wrapper.h且要求ENABLE_BF16编译宏与__CUDA_ARCH__ 800同时成立PyTorch 版将其替换为标准 CUDA 头cuda_bf16.h并把条件简化为仅__CUDA_ARCH__ 800。这正是 PyTorch 能同时支持 fp16 与 bf16 输入见下文算子约束的前提删除 MoE 相关与额外文件diff 中列出的Only in FasterTransformer...条目表明gemm_moe_problem_visitor.h、gemm_with_epilogue_visitor.h、moe_cutlass_kernel.h、moe_problem_visitor.h、compute_occupancy.h、epilogue_quant_helper.h及epilogue/threadblock子目录等 MoE 与推理服务栈相关的文件未被复制进一步印证 README 所说“只保留必要文件”。消费端_mixed_dtypes_linear算子的调用链这些扩展头文件在 PyTorch 中唯一的直接消费者是 MixedDtypesLinear.cu。该算子在 native_functions.yaml 中注册为torch._C._mixed_dtypes_linear内部算子- func: _mixed_dtypes_linear(Tensor input, Tensor weight, Tensor scale, *, Tensor? biasNone, str? activationNone) - Tensor dispatch: CUDA: _mixed_dtypes_linear从该实现可以梳理出完整的调用与约束链条平台与数据类型约束_mixed_dtypes_linear入口MixedDtypesLinear.cu仅在非 ROCm、非 Windows 构建下编译运行时要求 GPU 计算能力为8.xSM 80/86/89输入input必须为 fp16 或 bf16权重weight为uint8即 int8 权重打包为字节或QUInt4x24-bit 权重打包类型scale必须为与输入同 dtype 的 1D 张量bias可选为 1D 且与输入同 dtype。当weight.size(1) ! scale.size(0)时推断为 4-bit 量化QUInt4x2输入/权重要求 2D、strided、行主序行 stride 1 且列 stride 1输入的多维 batch 维会被压平为 2D权重形状要求行、列均能被 64 整除——这是 CUTLASS 混合精度 kernel 的硬限制length_k % 64 0 length_n % 64 0。CUTLASS kernel 的模板配置mixed_dtypes_linear_cutlassMixedDtypesLinear.cu这些参数与 README 所述“为支持的功能而保留的文件”一一对应配置项取值说明SmArchcutlass::arch::Sm80针对 Ampere 及以上架构编译ThreadblockShape32 × 128 × 64对应 ft_gemm_configs.h 中CtaShape32x128x64_WarpShape32x32x64WarpShape32 × 32 × 64warp 级 MMA 分块InstructionShape16 × 8 × 16Tensor Core 单条指令形状ThreadblockSwizzleGemmIdentityThreadblockSwizzle恒等 swizzle即 kernel 内部计算 grid 形状的基础OperatorOpMultiplyAddDequantizeInterleavedBToA来自 arch/mma.h 的 dequantize 乘加操作LayoutInputBColumnMajorTileInterleave64, 2K 维按ThreadblockK64分块交错列数 128B/4B ÷ 64 2直接使用了扩展库中的交错布局标签Stages4多级流水线缓冲数SplitKFactor1源码注释明确指出 1 会产生错误结果与 README diff 中删除 semaphore 的改动一致运行时调用链MixedDtypesLinear.cu构造Gemm::Arguments含 A/B/scale/bias/C/D 的TensorRef注意 B 的 leading dimension 乘以了kInterleave→gemm_op.can_implement(arguments)校验 →Gemm::get_workspace_size分配 workspace →initialize(...)绑定当前 CUDA stream →gemm_op.run(stream)发射 kernel →C10_CUDA_KERNEL_LAUNCH_CHECK()收尾。所有 CUTLASS 状态码经CUTLASS_STATUS_CHECK宏转换为TORCH_CHECK异常。bias/activation 的标签分发mixed_dtypes_linear_dispatch_bias_activation根据bias是否为空与activation字符串none/relu/silu选择fastertransformer::EpilogueOpNoBias、EpilogueOpBias、EpilogueOpBiasReLU、EpilogueOpBiasSilu四个标签最终实例化 epilogue_helpers.h 中对应的Epilogue特化——这正是扩展库 epilogue 层存在的意义。上游归宿等待 CUTLASS 官方收编README 的最后一段点明了这份代码的“临时性”根据 CUTLASS 项目方的讨论cutlass discussions #911 与 issues #1060CUTLASS 本身预期会原生包含这些扩展所支持的混合精度 GEMM 功能因此作者期望“这个目录最终会从 PyTorch 源码树中移除”。换言之cutlass_extensions是一份有明确退出计划的移植代码它填补了 CUTLASS 3.x 尚无 fp×int 混合精度 GEMM 官方支持的窗口期而 MixedDtypesLinear.cu 中的算子约束SM 8.x、维度 64 对齐、split-K 禁用也提示使用者这是一个针对特定量化推理场景的受限实现而非通用 dense linear 的替代路径。小结aten/src/ATen/native/cuda/cutlass_extensions的价值在于它以最小的改动面README 内附 diff 可逐行审计把 FasterTransformer 中成熟的“浮点激活 × 整型权重”GEMM 扩展层移植进 PyTorch并通过去掉 split-K workspace、放宽 BF16 条件、补齐invoke入口三处关键适配完成对 CUTLASS 3.x 的迁移。对阅读 PyTorch CUDA 后端源码的开发者而言这份目录连同 MixedDtypesLinear.cu 与 native_functions.yaml 中_mixed_dtypes_linear的注册构成了一条从算子 schema 到 CUTLASS kernel 模板实例化的、完整且可追踪的混合精度权重量化推理链路。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门