"""FP8 e4m3 x FP8 e4m3 GEMM for B200 (SM100) — real fp8 tensor-core MMA. y = (x @ w.T) * weight_scale -> bf16 x: (M, K) fp8_e4m3 row-major; w: (N, K) fp8_e4m3 row-major (K-major == the TN layout fp8 MMAs want); weight_scale: (N,) fp32 per-output-channel scale. Primary path: CUTLASS 4.x SM100 kernels (tcgen05 2SM UMMA, TMA warp-specialized, built via torch cpp_extension against the CUTLASS headers cloned into this directory). The per-channel scale is fused into the epilogue with a custom EVT (RowBroadcast(scale) * Accumulator), so the kernel writes scaled bf16 directly. Variants (picked per shape, measured on this box): * large shapes: MmaTile 256x256x128, cluster (2,2), 2SM MMA (~2.6-3.0 PFLOPS) * skinny M<=64 (decode): MmaTile 64x64x128, 1SM, maximizes CTA count for HBM streaming (the shape is memory-bound: must read all of W once) TMA requires 16B-aligned row strides, so for K=4127 both operands are zero-padded to K=4224 (identical result; zeros contribute nothing). The weight pad is cached keyed on (data_ptr, _version) so in-place weight updates (e.g. numeric-stress rescaling via copy_) invalidate it; x is re-padded on every call — nothing is memoized across differing inputs. Fallback path (if the extension can't build): a persistent Triton kernel using TMA tensor descriptors + tl.dot on fp8 (lowers to tcgen05) with warp specialization — same math, ~70% of the CUTLASS speed. """ import os from pathlib import Path # The PATH nvcc on this box is a broken wrapper (REAL_NVCC unset) and CUDA_HOME # may point at a toolkit that does not exist; force a real one. _ch = os.environ.get("CUDA_HOME") if not _ch or not os.path.exists(os.path.join(_ch, "bin", "nvcc")): for _cand in ("/usr/local/cuda-12.8", "/usr/local/cuda-12", "/usr/local/cuda"): if os.path.exists(os.path.join(_cand, "bin", "nvcc")): os.environ["CUDA_HOME"] = _cand break import torch import torch.nn as nn import triton import triton.language as tl E4M3_MAX = 448.0 _HERE = Path(__file__).resolve().parent # ========================================================================= # Triton fallback: persistent TMA-descriptor kernel (tcgen05 via tl.dot) # ========================================================================= def _alloc(size, alignment, stream): return torch.empty(size, device="cuda", dtype=torch.int8) triton.set_allocator(_alloc) def _configs(): return [ triton.Config({"BLOCK_M": 128, "BLOCK_N": 256, "BLOCK_K": 128, "GROUP_M": 8}, num_warps=8, num_stages=3), triton.Config({"BLOCK_M": 256, "BLOCK_N": 128, "BLOCK_K": 128, "GROUP_M": 16}, num_warps=8, num_stages=3), triton.Config({"BLOCK_M": 64, "BLOCK_N": 64, "BLOCK_K": 256, "GROUP_M": 1}, num_warps=4, num_stages=4), ] @triton.autotune(configs=_configs(), key=["M", "N", "K"]) @triton.jit def _fp8_gemm_desc_kernel( x_ptr, w_ptr, s_ptr, y_ptr, M, N, K, NUM_SMS, BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr, BLOCK_K: tl.constexpr, GROUP_M: tl.constexpr, ): num_pid_m = tl.cdiv(M, BLOCK_M) num_pid_n = tl.cdiv(N, BLOCK_N) k_tiles = tl.cdiv(K, BLOCK_K) num_tiles = num_pid_m * num_pid_n a_desc = tl.make_tensor_descriptor(x_ptr, [M, K], [K, 1], [BLOCK_M, BLOCK_K]) b_desc = tl.make_tensor_descriptor(w_ptr, [N, K], [K, 1], [BLOCK_N, BLOCK_K]) y_desc = tl.make_tensor_descriptor(y_ptr, [M, N], [N, 1], [BLOCK_M, BLOCK_N]) num_pid_in_group = GROUP_M * num_pid_n for tile_id in tl.range(tl.program_id(0), num_tiles, NUM_SMS, flatten=True, warp_specialize=True): group_id = tile_id // num_pid_in_group first_pid_m = group_id * GROUP_M group_size_m = tl.minimum(num_pid_m - first_pid_m, GROUP_M) pid_m = first_pid_m + ((tile_id % num_pid_in_group) % group_size_m) pid_n = (tile_id % num_pid_in_group) // group_size_m offs_am = pid_m * BLOCK_M offs_bn = pid_n * BLOCK_N acc = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) for ki in range(k_tiles): a = a_desc.load([offs_am, ki * BLOCK_K]) b = b_desc.load([offs_bn, ki * BLOCK_K]) acc = tl.dot(a, b.T, acc) offs_n = offs_bn + tl.arange(0, BLOCK_N) scale = tl.load(s_ptr + offs_n, mask=offs_n < N, other=0.0) out = (acc * scale[None, :]).to(tl.bfloat16) y_desc.store([offs_am, offs_bn], out) @triton.jit def _pad_k_kernel(src, dst, M, K, K_pad, BLOCK: tl.constexpr): """(M, K) contiguous fp8 -> (M, K_pad) contiguous, zero tail. Byte-wise.""" pid = tl.program_id(0) offs = pid * BLOCK + tl.arange(0, BLOCK) total = M * K_pad dmask = offs < total m = offs // K_pad k = offs - m * K_pad v = tl.load(src + m * K + k, mask=dmask & (k < K), other=0) tl.store(dst + offs, v, mask=dmask) def _pad_k(t8: torch.Tensor, K_pad: int) -> torch.Tensor: """Zero-pad the last dim of a contiguous fp8 (R, K) tensor to K_pad.""" R, K = t8.shape out = torch.empty((R, K_pad), device=t8.device, dtype=t8.dtype) total = R * K_pad _pad_k_kernel[(triton.cdiv(total, 4096),)]( t8.view(torch.uint8), out.view(torch.uint8), R, K, K_pad, BLOCK=4096, num_warps=4, ) return out # ========================================================================= # CUTLASS SM100 extension (primary path) # ========================================================================= _CUTLASS_SOURCE = r""" // SM100 (B200) fp8 e4m3 x e4m3 GEMM with per-column (N) scale epilogue. // y[m,n] = (sum_k x[m,k] * w[n,k]) * s[n], output bf16. // A = x (M,K) row-major e4m3; B = w (N,K) row-major == (K,N) column-major; // D = y (M,N) row-major bf16. Custom EVT: D = RowBroadcast(s) * Acc. #include #include #include "cutlass/cutlass.h" #include "cute/tensor.hpp" #include "cutlass/epilogue/collective/collective_builder.hpp" #include "cutlass/epilogue/fusion/sm90_visitor_tma_warpspecialized.hpp" #include "cutlass/gemm/collective/collective_builder.hpp" #include "cutlass/gemm/device/gemm_universal_adapter.h" #include "cutlass/gemm/kernel/gemm_universal.hpp" #include "cutlass/gemm/kernel/tile_scheduler.hpp" #include "cutlass/util/packed_stride.hpp" #include "cutlass/kernel_hardware_info.h" using namespace cute; #if !defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED) #error "CUTLASS_ARCH_MMA_SM100_SUPPORTED not defined - need CUDA 12.8+ and sm_100a" #endif using ElementA = cutlass::float_e4m3_t; using ElementB = cutlass::float_e4m3_t; using ElementD = cutlass::bfloat16_t; using ElementAcc = float; using ElementCompute = float; using LayoutA = cutlass::layout::RowMajor; // (M,K) using LayoutB = cutlass::layout::ColumnMajor; // (K,N) col-major == w (N,K) row-major using LayoutD = cutlass::layout::RowMajor; // (M,N) constexpr int AlignA = 16; constexpr int AlignB = 16; constexpr int AlignD = 8; constexpr auto RoundStyle = cutlass::FloatRoundStyle::round_to_nearest; template struct GemmConfig { using KernelSchedule = cute::conditional_t; using EpilogueSchedule = cute::conditional_t; // Per-CTA tile (2SM MMA splits M across the CTA pair) static constexpr int kTileM = decltype(size<0>(MmaTile{}))::value / (TwoSm ? 2 : 1); static constexpr int kTileN = decltype(size<1>(MmaTile{}))::value; static constexpr int kTileK = decltype(size<2>(MmaTile{}))::value; using CtaTileShape = Shape, Int, Int>; using RowScale = cutlass::epilogue::fusion::Sm90RowBroadcast< 0, CtaTileShape, float, float, Stride<_0, _1, _0>>; using CustomEVT = cutlass::epilogue::fusion::Sm90EVT< cutlass::epilogue::fusion::Sm90Compute, RowScale, cutlass::epilogue::fusion::Sm90AccFetch>; using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< cutlass::arch::Sm100, cutlass::arch::OpClassTensorOp, MmaTile, Cluster, cutlass::epilogue::collective::EpilogueTileAuto, ElementAcc, ElementCompute, void, LayoutD, AlignD, ElementD, LayoutD, AlignD, EpilogueSchedule, CustomEVT>::CollectiveOp; using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< cutlass::arch::Sm100, cutlass::arch::OpClassTensorOp, ElementA, LayoutA, AlignA, ElementB, LayoutB, AlignB, ElementAcc, MmaTile, Cluster, cutlass::gemm::collective::StageCountAutoCarveout< static_cast(sizeof(typename CollectiveEpilogue::SharedStorage))>, KernelSchedule>::CollectiveOp; using GemmKernel = cutlass::gemm::kernel::GemmUniversal< Shape, CollectiveMainloop, CollectiveEpilogue>; using Gemm = cutlass::gemm::device::GemmUniversalAdapter; }; template void run_gemm(torch::Tensor x, torch::Tensor w, torch::Tensor s, torch::Tensor y) { using Gemm = typename Config::Gemm; int M = x.size(0), K = x.size(1), N = w.size(0); using StrideA = typename Gemm::GemmKernel::StrideA; using StrideB = typename Gemm::GemmKernel::StrideB; using StrideD = typename Gemm::GemmKernel::StrideD; auto stride_A = cutlass::make_cute_packed_stride(StrideA{}, {M, K, 1}); auto stride_B = cutlass::make_cute_packed_stride(StrideB{}, {N, K, 1}); auto stride_D = cutlass::make_cute_packed_stride(StrideD{}, {M, N, 1}); static cutlass::KernelHardwareInfo hw_info = [] { cutlass::KernelHardwareInfo info; info.device_id = 0; info.sm_count = cutlass::KernelHardwareInfo::query_device_multiprocessor_count(0); return info; }(); typename Gemm::Arguments args{ cutlass::gemm::GemmUniversalMode::kGemm, {M, N, K, 1}, {reinterpret_cast(x.data_ptr()), stride_A, reinterpret_cast(w.data_ptr()), stride_B}, {{}, nullptr, stride_D, reinterpret_cast(y.data_ptr()), stride_D}, hw_info}; args.epilogue.thread = { {reinterpret_cast(s.data_ptr()), 0.f, {_0{}, _1{}, _0{}}}, // RowBroadcast(s) {}, // AccFetch {}, // multiplies }; Gemm gemm; auto status = gemm.can_implement(args); TORCH_CHECK(status == cutlass::Status::kSuccess, "CUTLASS cannot implement: ", cutlassGetStatusString(status), " M=", M, " N=", N, " K=", K); size_t ws_size = Gemm::get_workspace_size(args); void* ws_ptr = nullptr; torch::Tensor ws; if (ws_size > 0) { ws = torch::empty({static_cast(ws_size)}, torch::TensorOptions().dtype(torch::kUInt8).device(x.device())); ws_ptr = ws.data_ptr(); } auto stream = at::cuda::getCurrentCUDAStream().stream(); status = gemm.initialize(args, ws_ptr); TORCH_CHECK(status == cutlass::Status::kSuccess, "CUTLASS initialize failed"); status = gemm.run(stream); TORCH_CHECK(status == cutlass::Status::kSuccess, "CUTLASS run failed"); } // Large shapes: 2SM UMMA, MmaTile 256x256x128, cluster (2,2). using CfgBig = GemmConfig, Shape<_2, _2, _1>, true>; // Skinny decode shapes: 1SM, MmaTile 64x64x128 -> max CTA count for streaming W. using CfgSkinny = GemmConfig, Shape<_1, _1, _1>, false>; void fp8_gemm_big(torch::Tensor x, torch::Tensor w, torch::Tensor s, torch::Tensor y) { run_gemm(x, w, s, y); } void fp8_gemm_skinny(torch::Tensor x, torch::Tensor w, torch::Tensor s, torch::Tensor y) { run_gemm(x, w, s, y); } """ _EXT = None _EXT_TRIED = False def _cutlass_include_dirs(): root = _HERE / "cutlass" if not (root / "include" / "cutlass" / "cutlass.h").exists(): # Best-effort fetch if the checkout is missing (fresh workspace). import subprocess try: subprocess.run( ["git", "clone", "--depth", "1", "--branch", "v4.0.0", "https://github.com/NVIDIA/cutlass.git", str(root)], check=True, capture_output=True, timeout=600, ) except Exception: return None return [str(root / "include"), str(root / "tools" / "util" / "include")] def _get_ext(): global _EXT, _EXT_TRIED if _EXT is not None or _EXT_TRIED: return _EXT _EXT_TRIED = True try: incs = _cutlass_include_dirs() if incs is None: return None from torch.utils.cpp_extension import load_inline _EXT = load_inline( name="fp8_cutlass_sol", cpp_sources="", cuda_sources=_CUTLASS_SOURCE, extra_include_paths=incs, extra_cuda_cflags=[ "-O3", "-std=c++17", "-gencode=arch=compute_100a,code=sm_100a", "--expt-relaxed-constexpr", "-DNDEBUG", ], functions=["fp8_gemm_big", "fp8_gemm_skinny"], verbose=False, ) except Exception: _EXT = None return _EXT class Model(nn.Module): """y = ((x @ w.T) * weight_scale).to(bf16) with a genuine fp8 x fp8 MMA.""" def __init__(self, M: int, N: int, K: int): super().__init__() self.M, self.N, self.K = M, N, K w = torch.empty(N, K, dtype=torch.bfloat16) nn.init.normal_(w, std=0.02) s = (w.float().abs().amax(dim=1, keepdim=True) / E4M3_MAX).clamp(min=1e-12) w_fp8 = (w.float() / s).to(torch.float8_e4m3fn) self.register_buffer("weight", w_fp8) self.register_buffer("weight_scale", s.squeeze(1).to(torch.float32)) self._wpad = None # (key, padded weight) — invalidated on in-place update self._num_sms = None self._graphs = {} # (ptrs, versions) -> (CUDAGraph, output, warm count) self._graphs_ok = True def _padded_weight(self, K_pad: int) -> torch.Tensor: w = self.weight key = (w.data_ptr(), w._version, K_pad) if self._wpad is None or self._wpad[0] != key: self._wpad = (key, _pad_k(w.contiguous(), K_pad)) return self._wpad[1] def _compute(self, x: torch.Tensor) -> torch.Tensor: w = self.weight s = self.weight_scale M, K = x.shape N = w.shape[0] x = x.contiguous() if K % 16 != 0: # TMA needs 16B-aligned row strides: zero-pad K (identical result). K_pad = triton.cdiv(K, 128) * 128 x = _pad_k(x, K_pad) w = self._padded_weight(K_pad) K = K_pad if not w.is_contiguous(): w = w.contiguous() y = torch.empty((M, N), device=x.device, dtype=torch.bfloat16) ext = _get_ext() if ext is not None: try: if M <= 64: ext.fp8_gemm_skinny(x, w, s, y) else: ext.fp8_gemm_big(x, w, s, y) return y except Exception: pass # fall through to Triton if self._num_sms is None: self._num_sms = torch.cuda.get_device_properties(x.device).multi_processor_count grid = lambda META: ( min(self._num_sms, triton.cdiv(M, META["BLOCK_M"]) * triton.cdiv(N, META["BLOCK_N"])), ) _fp8_gemm_desc_kernel[grid](x, w, s, y, M, N, K, self._num_sms) return y def forward(self, x: torch.Tensor) -> torch.Tensor: # CUDA-graph the launch sequence (same treatment torch.compile # reduce-overhead gets). Keyed on input/weight identity AND version: # any new input tensor, in-place input write, or in-place weight # update changes the key and takes the eager path, so a replay never # serves stale data — the captured kernels re-read the live buffers. if self._graphs_ok and not torch.cuda.is_current_stream_capturing(): w = self.weight key = ( x.data_ptr(), x.shape, x._version, w.data_ptr(), w._version, self.weight_scale.data_ptr(), self.weight_scale._version, ) ent = self._graphs.get(key) if ent is not None and ent[0] is not None: ent[0].replay() return ent[1] count = 0 if ent is None else ent[2] if count >= 2: # Pipeline is warm (ext built, Triton autotune settled): capture. try: torch.cuda.synchronize() g = torch.cuda.CUDAGraph() with torch.cuda.graph(g): y = self._compute(x) self._graphs[key] = (g, y, count) g.replay() return y except Exception: self._graphs_ok = False else: self._graphs[key] = (None, None, count + 1) return self._compute(x) M = 4096 N = 4096 K = 4096 def get_inputs(): x = (torch.rand(M, K) * 8 - 4).to(torch.float8_e4m3fn) return [x] def get_init_inputs(): return [M, N, K]