From 1b937ebd3cb1fa05c5c49917ac2e32492541c9fd Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Fri, 21 Aug 2026 03:15:47 +0100 Subject: [PATCH 1/4] perf(ArcKernels): blockwise-FP8 GEMM on tensor cores, scale promoted per 128-K UNVERIFIED ON HARDWARE -- never run. No GPU was available for this wave. Nothing in this commit has executed. Every number below is a DERIVATION from published machine limits, not a measurement. Do not quote any of it as a result, and do not claim a speedup. WHAT `fp8_matmul_tiled` is 66% of a B=256 decode step -- 524 ms of 794. Its inner loop is `acc += s_input[ty][k] * s_weight[tx][k]`: two shared-memory float loads to feed one FMA, no register blocking, no tensor-core instruction anywhere in it. Shared memory delivers 32 floats/clk/SM against 128 FP32 lanes, so that loop cannot exceed ~1/4 of the FP32 CUDA-core rate however it is tuned. It is a structural ceiling, not a tuning problem. This adds `fp8_matmul_wmma` next to it: a real tensor-core GEMM that keeps the weights in FP8 and applies the `[N/128, K/128]` block scale by promoting an FP32 accumulator at every scale-block boundary along K -- DeepGEMM's structure. `K_BLK == block_size_x`, so one K-tile is exactly one scale block and the scale is a single scalar per tile. Numerically this is better than what it replaces, not a trade. The scalar kernel computes `f32(act) * f32(w * scale)` and rounds every product. Here the FP8 -> fp16/bf16 weight conversion is exact (e4m3's 3 mantissa bits and 2^-9..2^8 exponent range are representable in both), the tensor core forms each product at full width into f32, and the scale is applied once per 128 K-elements rather than once per element. DERIVED COST (a derivation, NOT a measurement) V4 at B=256: 7 sites/layer x 43 layers = 4.238 G params, so FLOP = 2 * 256 * 4.238e9 = 2.170 TFLOP. H200 FP32 non-tensor 67 TFLOP/s -> 32.4 ms (the scalar kernel's own bound; it achieves 6.8%) H200 BF16/FP16 tensor 989.5 TFLOP/s -> 2.19 ms <-- this kernel's bound H200 FP8 tensor 1979 TFLOP/s -> 1.10 ms HBM 4.8 TB/s over 4.238 GB -> 0.88 ms Compute-bound at a derived 2.19 ms floor against a measured 524 ms today. A WMMA GEMM of this shape typically realises 50-70% of peak, which would put it at a derived 3.1-4.4 ms. Arithmetic only; none of it has run. The remaining 2x to 1.10 ms is NOT reachable by tuning this kernel: `mma` with e4m3 operands needs BOTH operands in FP8, hence FP8 activations with their own per-token 128 scales. Named, not built. WHY NOT cuBLASLt, AND WHY NOT wgmma cuBLASLt does support exactly this layout -- `CUBLASLT_MATMUL_DESC_B_SCALE_MODE = CUBLASLT_MATMUL_MATRIX_SCALE_BLK128x128_32F` -- and the docs put it at compute capability 9.0, i.e. Hopper, our target. But the enum does not exist before CUDA 12.9 (absent in 12.4.1, absent in 12.8.0, present in 12.9.0) and .github/workflows/cuda_compile_check.yaml pins 12.4.1. It also wants FP8 activations with VEC128 scales and TN layout, so it is not the one-attribute change it first looks like. Checked before writing the kernel, as instructed. wgmma is reachable -- cudaforge auto-suffixes sm_90 to sm_90a -- and worth maybe another 1.3-1.5x. Not used, because this file cannot be run before it is committed and a hand-rolled wgmma descriptor that is subtly wrong returns plausible logits rather than an error. `nvcuda::wmma` fixes the fragment layout in the compiler, compiles for both sm_80 and sm_90a, and has a working precedent in this tree (kernels/mxfp4/mxfp4_gemm_wmma.cu). The wgmma rung belongs on a box that can run the A/B this commit sets up. KILL SWITCH Default-on per house fast-path-default policy, so the first box A/Bs both arms from one binary with no rebuild: (default) -> tensor-core WMMA GEMM ARC_NO_FP8_WMMA=1 -> the scalar `fp8_matmul_tiled`, i.e. exactly the behaviour of the commit before this one Run the control arm first. NO M THRESHOLD, DELIBERATELY This codebase has frozen a dispatch threshold from a single measured point twice -- `ARC_FP8_CUBLAS_MIN_M = 512` (512 was the only M measured, so M = 5..511 fell through to the scalar kernel) and `qtip::gather_policy`'s `n >= 683`. A threshold is a claim about every value it excludes. I have zero measured points, so inventing one here would be strictly worse than either. Sweep on hardware first, then add one if the sweep shows one. WHAT WAS ACTUALLY VERIFIED * `cargo check -p mistralrs-quant` (no cuda) green. * `cargo test -p mistralrs-quant --test cuda_kernel_build_guard` -- 5/5, including `expected_kernel_count_matches_disk`. * EXPECTED_KERNEL_COUNT 41 -> 43 (kernel + its cc<8.0 link stub); the dummy-stub count in that file's own note updated 5 -> 6. * nvcc for sm_80 and sm_90 via the free no-GPU GitHub Actions gate. Nothing else. In particular: no kernel executed, no output compared against a reference, no timing taken. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01UMmjFy8TvsgypxVWNVhhC7 --- mistralrs-quant/kernels/EXPECTED_KERNEL_COUNT | 4 +- .../blockwise_fp8/blockwise_fp8_gemm_wmma.cu | 528 ++++++++++++++++++ .../blockwise_fp8_gemm_wmma_dummy.cu | 61 ++ mistralrs-quant/src/blockwise_fp8/ffi.rs | 40 ++ mistralrs-quant/src/blockwise_fp8/ops.rs | 136 +++++ 5 files changed, 767 insertions(+), 2 deletions(-) create mode 100644 mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu create mode 100644 mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma_dummy.cu diff --git a/mistralrs-quant/kernels/EXPECTED_KERNEL_COUNT b/mistralrs-quant/kernels/EXPECTED_KERNEL_COUNT index b0a528a99..6cfd84858 100644 --- a/mistralrs-quant/kernels/EXPECTED_KERNEL_COUNT +++ b/mistralrs-quant/kernels/EXPECTED_KERNEL_COUNT @@ -25,7 +25,7 @@ # NOTE: this is the count the GLOB discovers, before cudaforge applies its # `exclude` patterns. It is therefore independent of compute capability. The # "Compiling N of M kernels" line in the build log shows the post-exclusion M -# (this count minus the 5 `*_dummy.cu` / `dummy_*.cu` stubs on an SM >= 80 +# (this count minus the 6 `*_dummy.cu` / `dummy_*.cu` stubs on an SM >= 80 # build), and that line is the engagement counter for a CUDA build: a CUDA # build that "succeeded" without it moving has not built what you think it did. -41 +43 diff --git a/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu b/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu new file mode 100644 index 000000000..9291ad74f --- /dev/null +++ b/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu @@ -0,0 +1,528 @@ +/** + * Parent system: ArcKernels. + * + * Blockwise-FP8 GEMM on TENSOR CORES, with an FP32 accumulator promoted at + * every scale-block boundary along K. + * + * ############################################################################ + * # UNVERIFIED ON HARDWARE -- never run. # + * # # + * # No GPU was available when this file was written. Nothing in it has # + * # executed. Every performance number in these comments is a DERIVATION # + * # from published machine limits, not a measurement, and is labelled as # + * # such. Do not quote any of them as a result. The first box that runs # + * # this must A/B it against `fp8_matmul_tiled` (see the kill switch note # + * # at the bottom) before anything here is believed. # + * ############################################################################ + * + * WHY THIS FILE EXISTS + * -------------------- + * `fp8_matmul_tiled` in blockwise_fp8_gemm.cu is a SCALAR CUDA-core GEMM. Its + * inner loop is + * + * acc += s_input[ty][k] * s_weight[tx][k]; + * + * -- two shared-memory float loads to feed ONE FMA, with no register blocking + * and no tensor-core instruction anywhere in it. Shared memory delivers 32 + * floats/clk/SM against 128 FP32 lanes, so that loop cannot exceed ~1/4 of the + * FP32 CUDA-core rate no matter how it is tuned; it is a structural ceiling, + * not a tuning problem. The fix is not a better scalar loop, it is to stop + * being scalar. + * + * THE ONE THING THAT MAKES BLOCKWISE FP8 AWKWARD + * ---------------------------------------------- + * The weight carries a `[ceil(N/bs_y), ceil(K/bs_x)]` grid of f32 scales + * (typically 128x128). The scale CHANGES every `bs_x` elements along the + * reduction axis, so you cannot hoist it out of the K loop, and you must not + * fold it into the operands: `bs` is an arbitrary f32 and rounding + * `fp8_value * scale` into bf16/fp16 to feed the tensor core would throw away + * most of the mantissa the scale is carrying. + * + * The structure that resolves this is DeepGEMM's: accumulate the UNSCALED + * products in an FP32 tensor-core accumulator across one scale block, then + * PROMOTE -- multiply that accumulator by the block's f32 scale and add it + * into a second, long-lived FP32 accumulator -- and reset. Hence + * `K_BLK == bs_x`: one K-tile is exactly one scale block, so the scale is a + * single scalar for the whole tile and the promotion happens once per tile. + * + * Numerically this is BETTER than the scalar kernel it replaces, not a + * trade. The scalar kernel computes `f32(act) * f32(w * scale)` and rounds + * every product to f32. Here the FP8 -> fp16/bf16 weight conversion is EXACT + * (e4m3 carries 3 mantissa bits and an exponent range of 2^-9..2^8; both fp16 + * and bf16 represent every one of the 256 e4m3 values exactly), the tensor + * core forms each product at full width into an f32 accumulator, and the + * scale is applied ONCE per 128 K-elements instead of once per element. + * + * WHY WMMA AND NOT wgmma + * ---------------------- + * `wgmma.mma_async` would be worth roughly another 1.3-1.5x on Hopper, and it + * is reachable here -- cudaforge auto-suffixes sm_90 to sm_90a + * (cudaforge-0.1.5 `compute_cap.rs::auto_suffix`: `b if b >= 90 => + * with_suffix(b, "a")`), which is the target wgmma requires. It is not used + * because this file cannot be executed before it is committed. Hand-rolled + * `wgmma` descriptors and swizzled shared-memory layouts are wrong far more + * often than they are right on the first try, and a wrong GEMM here is + * silent: it returns plausible logits. `nvcuda::wmma` fixes the fragment + * layout in the compiler, compiles unchanged for both sm_80 and sm_90a, and + * already has a working precedent in this tree (kernels/mxfp4/ + * mxfp4_gemm_wmma.cu). That precedent is the reason this is the version that + * ships first. The wgmma rung is a follow-on, and it should be written + * against a box that can run the A/B this one sets up. + * + * DERIVED COST (a derivation, NOT a measurement) + * ---------------------------------------------- + * V4 at B=256, the shape this targets: 7 blockwise-FP8 sites/layer x 43 + * layers = 4.238 G params, FLOP = 2 * 256 * 4.238e9 = 2.170 TFLOP. + * + * H200 FP32 non-tensor 67 TFLOP/s -> 32.4 ms (what the scalar + * kernel is bounded by, + * and it achieves ~6.8% + * of even that) + * H200 BF16/FP16 tensor 989.5 TFLOP/s -> 2.19 ms <-- THIS KERNEL'S BOUND + * H200 FP8 tensor 1979 TFLOP/s -> 1.10 ms + * HBM at 4.8 TB/s, 4.238 GB of FP8 -> 0.88 ms + * + * So this kernel is compute-bound at a DERIVED 2.19 ms floor, against a + * measured 524 ms for `fp8_matmul_tiled`. A WMMA GEMM of this shape typically + * realises 50-70% of the tensor-core peak, which would put it at a DERIVED + * 3.1-4.4 ms. All of that is arithmetic; none of it has run. + * + * The remaining 2x to the 1.10 ms FP8-tensor-core number is NOT reachable by + * tuning this kernel. `mma` with e4m3 operands requires BOTH operands in FP8, + * so it needs the ACTIVATION quantized to FP8 with its own per-token 128 + * scales -- a numerics change and a separate kernel, deliberately not done + * here. Named, not built. + * + * ALSO CHECKED, AND THE REASON THIS IS HAND-WRITTEN + * ------------------------------------------------- + * cuBLASLt gained exactly this layout -- `CUBLASLT_MATMUL_DESC_B_SCALE_MODE = + * CUBLASLT_MATMUL_MATRIX_SCALE_BLK128x128_32F` -- and it is documented at + * compute capability 9.0, i.e. Hopper, our target. But the enum does not + * exist before CUDA 12.9 (absent in 12.4.1, absent in 12.8.0, present in + * 12.9.0), and .github/workflows/cuda_compile_check.yaml pins 12.4.1. It also + * requires FP8 activations with VEC128 per-token scales and TN layout, so it + * is not the one-attribute change it first looks like. Worth revisiting if + * the toolkit pin ever moves; it is not available at the pin we build against. + * + * --use_fast_math (applied crate-wide by build.rs) -- WHICH SIDE WE ARE ON + * ----------------------------------------------------------------------- + * Unaffected. This kernel contains no transcendentals, no division and no + * sqrt, so the flags that rewrite those do not apply. `-fmad=true` only + * contracts the promotion `acc += c_frag.x[i] * ws` into an FMA, which is + * more accurate, not less. `--ftz=true` flushes f32 denormals, and nothing + * here goes near 1e-38: e4m3 spans 2^-9..2^8, the checkpoint scales are + * normal f32, and the accumulators are sums of such products. Tensor-core + * MMA is not governed by these flags at all. + */ + +#include +#include +#include +#include +#include + +#include +#include +#include +#include + +using namespace nvcuda::wmma; + +#define CUDA_CHECK(call) \ + do { \ + cudaError_t err = call; \ + if (err != cudaSuccess) { \ + fprintf(stderr, "CUDA error at %s:%d: %s\n", __FILE__, __LINE__, \ + cudaGetErrorString(err)); \ + } \ + } while (0) + +#define CEILDIV(x, y) (((x) + (y) - 1) / (y)) + +namespace fp8_gemm_wmma { + +// ============================================================================ +// Tiling +// ============================================================================ + +constexpr int WMMA_M_DIM = 16; +constexpr int WMMA_N_DIM = 16; +constexpr int WMMA_K_DIM = 16; + +// 8 warps: 4 along M, 2 along N; each warp owns two 16-wide N sub-tiles. +constexpr int WARPS_M = 4; +constexpr int WARPS_N = 2; +constexpr int N_SUBTILES = 2; +constexpr int WARPS_PER_BLOCK = WARPS_M * WARPS_N; // 8 +constexpr int BLOCK_THREADS = WARPS_PER_BLOCK * 32; // 256 + +constexpr int M_BLK = WARPS_M * WMMA_M_DIM; // 64 +constexpr int N_BLK = WARPS_N * N_SUBTILES * WMMA_N_DIM; // 64 + +// K_BLK is the scale-block length along K. The Rust dispatcher only selects +// this kernel when `block_size_x % K_BLK == 0`, which is what makes the scale +// a single scalar per (block, k-tile) and lets the promotion happen once per +// tile. Changing this constant changes that contract -- update +// `fp8_wmma_eligible` in blockwise_fp8/ops.rs with it. +constexpr int K_BLK = 128; +constexpr int WMMA_K_STEPS = K_BLK / WMMA_K_DIM; // 8 + +// Shared-memory row padding, in ELEMENTS of T. +// +// 8, and it must stay a multiple of 8. `wmma::load_matrix_sync` requires the +// leading dimension to be a multiple of 16 bytes for 16-bit element types; +// (128 + 8) * 2 B = 272 B = 17 * 16 B, which satisfies it. This is the +// opposite of the right answer for the SCALAR kernel next door, where a +1 +// pad breaks a 4-way bank conflict -- there the loads are ordinary +// per-element `LDS.32` with no alignment rule, here they are fragment loads +// that have one. A +1 pad here would make the leading dimension 258 B, which +// is not a multiple of 16 B, and the load would be undefined. Do not +// "simplify" this to +1 to match the other kernel. +constexpr int SMEM_PAD = 8; +constexpr int A_LDM = K_BLK + SMEM_PAD; // 136 elements +constexpr int B_LDM = K_BLK + SMEM_PAD; // 136 elements + +// One K-tile of A and of B, plus (aliased over A) the f32 output staging tile. +constexpr int A_SMEM_ELEMS = M_BLK * A_LDM; // 8704 +constexpr int B_SMEM_ELEMS = N_BLK * B_LDM; // 8704 +constexpr int C_SMEM_ELEMS = M_BLK * N_BLK; // 4096 floats + +// Bytes, for a 16-bit T. Static shared memory, so we stay under the 48 KB +// per-block limit that needs no `cudaFuncSetAttribute` opt-in: +// (8704 + 8704) * 2 = 34,816 B. +// The f32 staging tile is 4096 * 4 = 16,384 B and ALIASES A_sh (17,408 B), +// which is only safe because it is written after the K loop has finished with +// A. The __syncthreads() before that reuse is load-bearing. +constexpr int AB_SMEM_BYTES = (A_SMEM_ELEMS + B_SMEM_ELEMS) * 2; +static_assert(C_SMEM_ELEMS * sizeof(float) <= A_SMEM_ELEMS * 2, + "f32 output staging tile must fit inside the A tile it aliases"); +static_assert(AB_SMEM_BYTES <= 48 * 1024, + "static shared memory must stay under the 48 KB no-opt-in limit"); + +using AccFrag = fragment; + +// Number of f32 registers a 16x16 accumulator fragment exposes per lane. +// Asserted rather than assumed: if a future CUDA changes the fragment layout, +// this fails the build instead of silently promoting the wrong registers. +constexpr int ACC_REGS = 8; + +// ============================================================================ +// FP8 -> T conversion +// +// Built on `__nv_cvt_fp8_to_halfraw`, the same primitive the scalar kernel in +// blockwise_fp8_gemm.cu already uses, so the conversion path is one this tree +// has compiled before. It is a __host__ __device__ function from cuda_fp8.h +// and is available regardless of compute capability -- on pre-sm_89 it is +// emulated in software rather than absent, which is what lets this file +// compile and be correct for the sm_80 CI lane as well as sm_90a. +// +// EXACTNESS: e4m3 has 3 mantissa bits and unbiased exponents in [-9, 8]. fp16 +// (10 mantissa bits, exp [-24, 15]) and bf16 (7 mantissa bits, exp +// [-126, 127]) each represent all 256 e4m3 values exactly, so no scale +// information is lost here. The scale itself is deliberately NOT applied at +// this point -- it is applied to the f32 accumulator in the promotion step. +// ============================================================================ + +__device__ __forceinline__ half fp8_to_half(uint8_t bits) { + // Deliberately the same expression shape the scalar kernel next door + // already compiles: `__nv_cvt_fp8_to_halfraw` returns `__half_raw`, and the + // implicit `__half_raw` -> `__half` conversion is available because + // build.rs passes `-U__CUDA_NO_HALF_CONVERSIONS__`. + const __half_raw hr = __nv_cvt_fp8_to_halfraw(bits, __NV_E4M3); + return hr; +} + +template __device__ __forceinline__ T fp8_to_T(uint8_t bits); + +template <> __device__ __forceinline__ half fp8_to_T(uint8_t bits) { + return fp8_to_half(bits); +} + +template <> +__device__ __forceinline__ __nv_bfloat16 fp8_to_T<__nv_bfloat16>(uint8_t bits) { + return __float2bfloat16(__half2float(fp8_to_half(bits))); +} + +template __device__ __forceinline__ T zero_of(); +template <> __device__ __forceinline__ half zero_of() { + return __float2half(0.0f); +} +template <> __device__ __forceinline__ __nv_bfloat16 zero_of<__nv_bfloat16>() { + return __float2bfloat16(0.0f); +} + +template __device__ __forceinline__ float T_to_float(T v); +template <> __device__ __forceinline__ float T_to_float(half v) { + return __half2float(v); +} +template <> +__device__ __forceinline__ float T_to_float<__nv_bfloat16>(__nv_bfloat16 v) { + return __bfloat162float(v); +} + +template __device__ __forceinline__ T float_to_T(float v); +template <> __device__ __forceinline__ half float_to_T(float v) { + return __float2half(v); +} +template <> +__device__ __forceinline__ __nv_bfloat16 float_to_T<__nv_bfloat16>(float v) { + return __float2bfloat16(v); +} + +// ============================================================================ +// The kernel +// +// C[M, N] = A[M, K] * B[N, K]^T, A in T, B in FP8 e4m3 with a +// [ceil(N/bs_y), ceil(K/bs_x)] f32 scale grid, C in T. +// +// Grid: (CEILDIV(N, N_BLK), CEILDIV(M, M_BLK)). Block: 256 threads. +// +// PRECONDITIONS, enforced by `fp8_wmma_eligible` in blockwise_fp8/ops.rs: +// * block_size_y % N_BLK == 0 (a 64-wide, 64-aligned N tile never +// straddles a scale row) +// * block_size_x % K_BLK == 0 (a K tile is exactly one scale block, so +// the scale is one scalar per tile) +// M, N and K are otherwise unconstrained; ragged edges are zero-filled. +// ============================================================================ + +template +__launch_bounds__(BLOCK_THREADS) __global__ void fp8_matmul_wmma( + const T *__restrict__ input, // [M, K] + const __nv_fp8_e4m3 *__restrict__ weight, // [N, K] row-major + const float *__restrict__ weight_scale, // [ceil(N/bs_y), ceil(K/bs_x)] + T *__restrict__ output, // [M, N] + int M, int N, int K, int scale_row_stride, int block_size_y, + int block_size_x) { + static_assert(AccFrag::num_elements == ACC_REGS, + "16x16x16 f32 accumulator fragment is not 8 registers/lane; " + "the promotion loop below assumes it is"); + + __shared__ __align__(16) char smem_raw[AB_SMEM_BYTES]; + T *A_sh = reinterpret_cast(smem_raw); + T *B_sh = reinterpret_cast(smem_raw) + A_SMEM_ELEMS; + + const int threadId = threadIdx.x; + const int warpId = threadId >> 5; + const int warp_m_idx = warpId / WARPS_N; // 0..3 + const int warp_n_idx = warpId % WARPS_N; // 0..1 + + const int m_base = blockIdx.y * M_BLK; + const int n_base = blockIdx.x * N_BLK; + + // The scale row is fixed for the whole block: the precondition + // `block_size_y % N_BLK == 0` guarantees rows [n_base, n_base + N_BLK) all + // land in scale row n_base / block_size_y. + const int scale_row_off = (n_base / block_size_y) * scale_row_stride; + + // Long-lived FP32 accumulators. `c_frag` accumulates ONE scale block and is + // reset after every promotion; `acc` carries the scaled running total for + // the whole K extent. + AccFrag c_frag[N_SUBTILES]; + float acc[N_SUBTILES][ACC_REGS]; +#pragma unroll + for (int s = 0; s < N_SUBTILES; ++s) { + fill_fragment(c_frag[s], 0.0f); +#pragma unroll + for (int i = 0; i < ACC_REGS; ++i) + acc[s][i] = 0.0f; + } + + const T kZero = zero_of(); + + for (int k_base = 0; k_base < K; k_base += K_BLK) { + // ---- Stage A[M_BLK, K_BLK] ------------------------------------------ + // 64 * 128 = 8192 elements over 256 threads = 32 each. Loaded 8 at a time + // (16 B) when the row is fully in range, which is the common case. + constexpr int A_ELEMS = M_BLK * K_BLK; + constexpr int VEC = 8; // 8 * 2 B = 16 B + for (int i = threadId * VEC; i < A_ELEMS; i += BLOCK_THREADS * VEC) { + const int lm = i / K_BLK; + const int lk = i % K_BLK; + const int gm = m_base + lm; + const int gk = k_base + lk; + T *dst = &A_sh[lm * A_LDM + lk]; + + if (gm < M && gk + VEC <= K) { + // A row of `input` is K elements of 2 B; the 16 B load is aligned + // whenever K % 8 == 0, which the dispatcher does not require, so go + // through a bytewise-safe vector type only when the address allows. + const T *src = &input[(size_t)gm * K + gk]; + if ((reinterpret_cast(src) & 0xF) == 0) { + *reinterpret_cast(dst) = + *reinterpret_cast(src); + } else { +#pragma unroll + for (int v = 0; v < VEC; ++v) + dst[v] = src[v]; + } + } else { +#pragma unroll + for (int v = 0; v < VEC; ++v) { + const int gk_v = gk + v; + dst[v] = (gm < M && gk_v < K) ? input[(size_t)gm * K + gk_v] : kZero; + } + } + } + + // ---- Stage B[N_BLK, K_BLK], converted FP8 -> T, UNSCALED ------------- + // 64 * 128 = 8192 FP8 bytes over 256 threads = 32 each, read as 8 + // uint32_t (4 weights per load). + constexpr int B_ELEMS = N_BLK * K_BLK; + constexpr int WVEC = 4; // 4 * 1 B = 4 B + for (int i = threadId * WVEC; i < B_ELEMS; i += BLOCK_THREADS * WVEC) { + const int ln = i / K_BLK; + const int lk = i % K_BLK; + const int gn = n_base + ln; + const int gk = k_base + lk; + T *dst = &B_sh[ln * B_LDM + lk]; + + if (gn < N && gk + WVEC <= K) { + const uint8_t *src = + reinterpret_cast(&weight[(size_t)gn * K + gk]); + if ((reinterpret_cast(src) & 0x3) == 0) { + const uint32_t w4 = __ldg(reinterpret_cast(src)); +#pragma unroll + for (int v = 0; v < WVEC; ++v) + dst[v] = fp8_to_T(static_cast((w4 >> (8 * v)) & 0xFF)); + } else { +#pragma unroll + for (int v = 0; v < WVEC; ++v) + dst[v] = fp8_to_T(__ldg(src + v)); + } + } else { +#pragma unroll + for (int v = 0; v < WVEC; ++v) { + const int gk_v = gk + v; + if (gn < N && gk_v < K) { + dst[v] = fp8_to_T(__ldg(reinterpret_cast( + &weight[(size_t)gn * K + gk_v]))); + } else { + // Zero weight, so the padded lanes contribute nothing to the + // accumulator and the promotion below stays correct on ragged + // edges without any masking. + dst[v] = kZero; + } + } + } + } + + __syncthreads(); + + // ---- Tensor-core accumulation over this one scale block -------------- +#pragma unroll + for (int k_step = 0; k_step < WMMA_K_STEPS; ++k_step) { + fragment + a_frag; + load_matrix_sync(a_frag, + A_sh + warp_m_idx * WMMA_M_DIM * A_LDM + + k_step * WMMA_K_DIM, + A_LDM); + +#pragma unroll + for (int s = 0; s < N_SUBTILES; ++s) { + // B_sh is [N][K]; as a K x N operand that is column-major with + // leading dimension B_LDM. + fragment + b_frag; + load_matrix_sync(b_frag, + B_sh + + (warp_n_idx * N_SUBTILES + s) * WMMA_N_DIM * B_LDM + + k_step * WMMA_K_DIM, + B_LDM); + mma_sync(c_frag[s], a_frag, b_frag, c_frag[s]); + } + } + + // ---- PROMOTE: apply this block's f32 scale, fold in, reset ----------- + // This is the whole reason the kernel is shaped this way. `ws` is one + // scalar for the entire tile because K_BLK divides block_size_x and + // N_BLK divides block_size_y. + { + const float ws = __ldg(&weight_scale[scale_row_off + k_base / block_size_x]); +#pragma unroll + for (int s = 0; s < N_SUBTILES; ++s) { +#pragma unroll + for (int i = 0; i < ACC_REGS; ++i) + acc[s][i] += c_frag[s].x[i] * ws; + fill_fragment(c_frag[s], 0.0f); + } + } + + __syncthreads(); + } + + // ---- Store ------------------------------------------------------------ + // Move the promoted f32 totals back into fragments so store_matrix_sync can + // lay them out, then stage in shared memory and write coalesced. + // + // C_sh ALIASES A_sh. The __syncthreads() at the end of the K loop above is + // what makes that safe -- every warp is done reading A before any warp + // writes C over it. + float *C_sh = reinterpret_cast(smem_raw); +#pragma unroll + for (int s = 0; s < N_SUBTILES; ++s) { +#pragma unroll + for (int i = 0; i < ACC_REGS; ++i) + c_frag[s].x[i] = acc[s][i]; + store_matrix_sync(C_sh + warp_m_idx * WMMA_M_DIM * N_BLK + + (warp_n_idx * N_SUBTILES + s) * WMMA_N_DIM, + c_frag[s], N_BLK, mem_row_major); + } + __syncthreads(); + + for (int i = threadId; i < C_SMEM_ELEMS; i += BLOCK_THREADS) { + const int lm = i / N_BLK; + const int ln = i % N_BLK; + const int gm = m_base + lm; + const int gn = n_base + ln; + if (gm < M && gn < N) + output[(size_t)gm * N + gn] = float_to_T(C_sh[lm * N_BLK + ln]); + } +} + +} // namespace fp8_gemm_wmma + +// ============================================================================ +// C API +// +// Signature-identical to `launch_fp8_matmul_{f16,bf16}` in +// blockwise_fp8_gemm.cu, so the Rust dispatcher can swap one for the other +// with no other change. That is deliberate: it is what makes the kill switch +// a single boolean. +// ============================================================================ + +extern "C" void launch_fp8_matmul_wmma_f16( + const __half *input, const __nv_fp8_e4m3 *weight, const float *weight_scale, + __half *output, int M, int N, int K, int scale_row_stride, int block_size_y, + int block_size_x, cudaStream_t stream) { + dim3 block(fp8_gemm_wmma::BLOCK_THREADS); + dim3 grid(CEILDIV(N, fp8_gemm_wmma::N_BLK), CEILDIV(M, fp8_gemm_wmma::M_BLK)); + + fp8_gemm_wmma::fp8_matmul_wmma<<>>( + input, weight, weight_scale, output, M, N, K, scale_row_stride, + block_size_y, block_size_x); + CUDA_CHECK(cudaGetLastError()); +} + +extern "C" void launch_fp8_matmul_wmma_bf16( + const __nv_bfloat16 *input, const __nv_fp8_e4m3 *weight, + const float *weight_scale, __nv_bfloat16 *output, int M, int N, int K, + int scale_row_stride, int block_size_y, int block_size_x, + cudaStream_t stream) { + dim3 block(fp8_gemm_wmma::BLOCK_THREADS); + dim3 grid(CEILDIV(N, fp8_gemm_wmma::N_BLK), CEILDIV(M, fp8_gemm_wmma::M_BLK)); + + fp8_gemm_wmma::fp8_matmul_wmma<__nv_bfloat16><<>>( + input, weight, weight_scale, output, M, N, K, scale_row_stride, + block_size_y, block_size_x); + CUDA_CHECK(cudaGetLastError()); +} + +// Reported to the Rust side so the dispatcher's eligibility test cannot drift +// from the kernel's actual tiling. See `fp8_wmma_eligible` in +// blockwise_fp8/ops.rs. +extern "C" void fp8_matmul_wmma_tile_dims(int *m_blk, int *n_blk, int *k_blk) { + *m_blk = fp8_gemm_wmma::M_BLK; + *n_blk = fp8_gemm_wmma::N_BLK; + *k_blk = fp8_gemm_wmma::K_BLK; +} diff --git a/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma_dummy.cu b/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma_dummy.cu new file mode 100644 index 000000000..19a443895 --- /dev/null +++ b/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma_dummy.cu @@ -0,0 +1,61 @@ +/** + * Parent system: ArcKernels. + * + * Link-only stubs for the WMMA blockwise-FP8 GEMM on compute capability < 8.0. + * + * UNVERIFIED ON HARDWARE -- never run. (Nothing in this wave has executed.) + * + * WHY THIS FILE EXISTS + * -------------------- + * `build.rs` excludes `*_wmma.cu` from the kernel set when the compute cap is + * below 8.0 -- BF16 WMMA fragments need sm_80 -- and excludes `*_dummy.cu` + * when it is 8.0 or above. So exactly one of this file and + * blockwise_fp8_gemm_wmma.cu is compiled, and the `extern "C"` symbols the + * Rust FFI declares resolve either way. + * + * Nothing here is ever reached at runtime: `has_blockwise_fp8_kernels` is + * only emitted for cc >= 8.0, so `HAVE_BLOCKWISE_GEMM_KERNELS` is false on + * the builds that link these stubs and the Rust dispatcher never calls them. + * They exist because an rlib build does not link, so a missing symbol here + * would surface only when a binary is finally linked -- on a rented box. + * + * Mirrors blockwise_fp8_gemm_dummy.cu, which does the same job for the scalar + * kernels next door. + */ + +#include +#include +#include +#include +#include +#include + +extern "C" void launch_fp8_matmul_wmma_f16(const __half *input, + const void *weight, // __nv_fp8_e4m3* + const float *weight_scale, + __half *output, int M, int N, int K, + int scale_row_stride, + int block_size_y, int block_size_x, + cudaStream_t stream) { + fprintf(stderr, "FP8 WMMA matmul not supported on this GPU (requires compute " + "capability >= 8.0)\n"); +} + +extern "C" void launch_fp8_matmul_wmma_bf16( + const __nv_bfloat16 *input, + const void *weight, // __nv_fp8_e4m3* + const float *weight_scale, __nv_bfloat16 *output, int M, int N, int K, + int scale_row_stride, int block_size_y, int block_size_x, + cudaStream_t stream) { + fprintf(stderr, "FP8 WMMA matmul not supported on this GPU (requires compute " + "capability >= 8.0)\n"); +} + +// Must still report the tiling the Rust dispatcher asks about at startup. The +// values are the ones in blockwise_fp8_gemm_wmma.cu; they are only used to +// decide eligibility, and on this build the kernel is never selected anyway. +extern "C" void fp8_matmul_wmma_tile_dims(int *m_blk, int *n_blk, int *k_blk) { + *m_blk = 64; + *n_blk = 64; + *k_blk = 128; +} diff --git a/mistralrs-quant/src/blockwise_fp8/ffi.rs b/mistralrs-quant/src/blockwise_fp8/ffi.rs index 4302fd785..0596a83df 100644 --- a/mistralrs-quant/src/blockwise_fp8/ffi.rs +++ b/mistralrs-quant/src/blockwise_fp8/ffi.rs @@ -113,6 +113,46 @@ extern "C" { stream: candle_core::cuda::cudarc::driver::sys::CUstream, ); + // ---- Tensor-core blockwise-FP8 GEMM (kernels/blockwise_fp8/ + // blockwise_fp8_gemm_wmma.cu). Signature-identical to + // `launch_fp8_matmul_*` above so the two are swappable behind one boolean. + // + // UNVERIFIED ON HARDWARE -- never run. See the kernel's header comment. + pub(crate) fn launch_fp8_matmul_wmma_f16( + input: *const f16, + weight: *const F8E4M3, + weight_scale: *const f32, + output: *mut f16, + m: i32, + n: i32, + k: i32, + scale_row_stride: i32, + block_size_y: i32, + block_size_x: i32, + stream: candle_core::cuda::cudarc::driver::sys::CUstream, + ); + + pub(crate) fn launch_fp8_matmul_wmma_bf16( + input: *const bf16, + weight: *const F8E4M3, + weight_scale: *const f32, + output: *mut bf16, + m: i32, + n: i32, + k: i32, + scale_row_stride: i32, + block_size_y: i32, + block_size_x: i32, + stream: candle_core::cuda::cudarc::driver::sys::CUstream, + ); + + /// Report the WMMA kernel's block tiling so the Rust eligibility test + /// cannot drift from the tiling the kernel was actually compiled with. + /// The dispatcher's preconditions (`block_size_y % n_blk == 0`, + /// `block_size_x % k_blk == 0`) are what make the block scale a single + /// scalar per tile, which is the kernel's whole premise. + pub(crate) fn fp8_matmul_wmma_tile_dims(m_blk: *mut i32, n_blk: *mut i32, k_blk: *mut i32); + // FP8 GEMV kernels (dedicated decode path, M <= 4). Warp-per-row, // dequant-in-registers with per-block scales, f32 accumulate. pub(crate) fn launch_fp8_gemv_f16( diff --git a/mistralrs-quant/src/blockwise_fp8/ops.rs b/mistralrs-quant/src/blockwise_fp8/ops.rs index 85eb12b23..c75bfb249 100644 --- a/mistralrs-quant/src/blockwise_fp8/ops.rs +++ b/mistralrs-quant/src/blockwise_fp8/ops.rs @@ -925,6 +925,95 @@ fn fp8_gemv_max_m() -> usize { *MAX_M } +/// Parent system: ArcKernels. +/// +/// Kill switch for the tensor-core blockwise-FP8 GEMM +/// (`kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu`). +/// +/// # ๐Ÿ”ด UNVERIFIED ON HARDWARE โ€” never run +/// +/// The WMMA kernel has never executed. It is nonetheless the DEFAULT, per the +/// house fast-path-default policy, so that the first box to run this binary +/// A/Bs it against the scalar kernel without a rebuild: +/// +/// ```text +/// (default) -> tensor-core WMMA GEMM +/// ARC_NO_FP8_WMMA=1 -> the shipped scalar `fp8_matmul_tiled` +/// ``` +/// +/// Set `ARC_NO_FP8_WMMA=1` to get back exactly the behaviour of the commit +/// before this one. That is the control arm; run it first. +/// +/// # Why there is no M threshold here +/// +/// There is an obvious temptation to gate this on some minimum M, since a +/// 64-row block tile is mostly empty at M = 8. Two reasons not to, and the +/// second is the important one: +/// +/// 1. In that regime the kernel is bound by streaming the weight tile, not by +/// the MMA issue rate, so an under-filled M tile costs far less than the +/// ~100x per-useful-FLOP penalty of staying on the scalar kernel. +/// 2. **This codebase has now frozen a dispatch threshold from a single +/// measured point twice** โ€” `ARC_FP8_CUBLAS_MIN_M = 512` (512 was the only +/// M ever measured, and M = 5..511 fell through to the scalar kernel as a +/// result) and `qtip::gather_policy`'s `n >= 683` tile-fill predicate. A +/// threshold is a claim about every value it excludes. I have **zero** +/// measured points, so inventing one here would be strictly worse than +/// those two were. Sweep it on hardware, then add a threshold if the sweep +/// shows one โ€” with the clean rows written down. +#[cfg(feature = "cuda")] +fn fp8_wmma_enabled() -> bool { + use std::sync::LazyLock; + static ENABLED: LazyLock = LazyLock::new(|| std::env::var("ARC_NO_FP8_WMMA").is_err()); + *ENABLED +} + +/// Block tiling the WMMA kernel was actually compiled with, read from the +/// kernel itself rather than duplicated here. +/// +/// The eligibility test below is a statement about that tiling: it is what +/// guarantees the `[N/bs_y, K/bs_x]` block scale is a single scalar over each +/// tile, which is the kernel's entire premise. Duplicating `64` and `128` in +/// Rust would let the two drift silently the day someone retunes the kernel, +/// and the failure mode of that drift is a *numerically wrong* GEMM, not a +/// crash โ€” so the numbers are exported from the `.cu` and read once. +#[cfg(feature = "cuda")] +fn fp8_wmma_tile_dims() -> (i32, i32, i32) { + use std::sync::LazyLock; + static DIMS: LazyLock<(i32, i32, i32)> = LazyLock::new(|| { + let (mut m, mut n, mut k) = (0i32, 0i32, 0i32); + // Pure host function; writes three ints. Linked from either + // blockwise_fp8_gemm_wmma.cu or its `_dummy` stub. + unsafe { crate::blockwise_fp8::ffi::fp8_matmul_wmma_tile_dims(&mut m, &mut n, &mut k) }; + (m, n, k) + }); + *DIMS +} + +/// Whether a given blockwise-FP8 scale geometry is one the WMMA kernel can +/// serve correctly. +/// +/// The kernel applies one f32 scale to the whole `N_BLK x K_BLK` accumulator +/// tile per K step. That is only correct when a tile cannot straddle two +/// scale blocks: +/// +/// * `block_size_x % K_BLK == 0` โ€” a K tile lies inside one scale column. +/// * `block_size_y % N_BLK == 0` โ€” an N tile lies inside one scale row (N +/// tiles are `N_BLK`-aligned, so divisibility is sufficient). +/// +/// For the usual DeepSeek-style `[128, 128]` geometry, with `N_BLK = 64` and +/// `K_BLK = 128`, both hold. Anything else falls back to the scalar kernel, +/// which reads the scale per element and so has no such constraint. M, N and +/// K themselves are unconstrained โ€” the kernel zero-fills ragged edges. +#[cfg(feature = "cuda")] +fn fp8_wmma_eligible(block_size_y: i32, block_size_x: i32) -> bool { + let (_m_blk, n_blk, k_blk) = fp8_wmma_tile_dims(); + if n_blk <= 0 || k_blk <= 0 { + return false; + } + block_size_y % n_blk == 0 && block_size_x % k_blk == 0 +} + /// FP8 blockwise matmul. /// Computes output = input @ weight.T where weight is FP8 blockwise quantized. /// - input: [M, K] in fp16/bf16 @@ -1012,6 +1101,19 @@ fn fp8_blockwise_matmul_impl( let gemv_aligned = k % 4 == 0 && block_size_x % 4 == 0; let use_gemv = force_gemv.unwrap_or_else(|| (m as usize) <= fp8_gemv_max_m()) && gemv_aligned; + // Tensor-core GEMM for everything the GEMV does not own. + // + // ๐Ÿ”ด UNVERIFIED ON HARDWARE โ€” never run. Default-on so the first box can + // A/B it in one binary; `ARC_NO_FP8_WMMA=1` restores the scalar kernel. + // + // `force_gemv` keeps its documented meaning: an explicit `Some(false)` + // still selects the *scalar* tiled GEMM, so any test that pins that path + // keeps testing it. + let use_wmma = !use_gemv + && force_gemv.is_none() + && fp8_wmma_enabled() + && fp8_wmma_eligible(block_size_y, block_size_x); + let input_l = input.layout(); let weight_l = weight.layout(); let scales_l = scales.layout(); @@ -1062,6 +1164,23 @@ fn fp8_blockwise_matmul_impl( dev.cuda_stream().cu_stream(), ) }; + } else if use_wmma { + // UNVERIFIED ON HARDWARE โ€” never run. + unsafe { + ffi::launch_fp8_matmul_wmma_f16( + input_ptr as *const _, + weight_ptr as *const _, + scales_ptr as *const _, + output_ptr as *mut _, + m, + n, + k, + scale_row_stride, + block_size_y, + block_size_x, + dev.cuda_stream().cu_stream(), + ) + }; } else { unsafe { ffi::launch_fp8_matmul_f16( @@ -1116,6 +1235,23 @@ fn fp8_blockwise_matmul_impl( dev.cuda_stream().cu_stream(), ) }; + } else if use_wmma { + // UNVERIFIED ON HARDWARE โ€” never run. + unsafe { + ffi::launch_fp8_matmul_wmma_bf16( + input_ptr as *const _, + weight_ptr as *const _, + scales_ptr as *const _, + output_ptr as *mut _, + m, + n, + k, + scale_row_stride, + block_size_y, + block_size_x, + dev.cuda_stream().cu_stream(), + ) + }; } else { unsafe { ffi::launch_fp8_matmul_bf16( From 776eafabb0770da769916d71d31d9aecec27c0b2 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Fri, 21 Aug 2026 03:19:46 +0100 Subject: [PATCH 2/4] fix(ArcKernels): WMMA ldm must be a multiple of 32 B, not 16 -- the C++ guide understates the ISA MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit UNVERIFIED ON HARDWARE -- never run. Two NVIDIA documents govern `wmma::load_matrix_sync`'s leading dimension and they DISAGREE. The laxer one is the one people quote, and the first version of this kernel followed it -- which compiles, and is wrong. CUDA C++ Programming Guide 12.4 ยง7.24.1: "mptr must be a 256-bit aligned pointer ... and [ldm] must be a multiple of 8 for __half element type or multiple of 4 for float element type. (i.e., multiple of 16 bytes in both cases)." PTX ISA 12.4 ยง9.7.13.3.2, working our exact shape through: "The starting address of each instance of the leading dimension (row or column) must be aligned with the size of the corresponding fragment in bytes." ... for `wmma.load.a.sync.aligned.row.m16n16k16.f16` the fragment is 32 B (eight `.f16x2` elements), so "p is a multiple of 32" and "2*s is a multiple of 32" i.e. ldm must be a multiple of SIXTEEN __half elements, double the guide's number, and the base pointer must be 32 B aligned, not 16. The +8 pad gave ldm = 136 elements = 272 B. 272 % 16 == 0 satisfies the guide; 272 % 32 == 16 violates the ISA. Worst of both worlds, because nothing would have complained until the numbers came out subtly wrong on a rented box. * SMEM_PAD 8 -> 16, so ldm = 144 elements = 288 B = 9 * 32 B. Satisfies both documents, for f16 and bf16 alike. Asserted at compile time now, with the ISA rule as the assertion message so the next person does not "optimise" the pad back down to the guide's number. * The shared tile was `__align__(16)`; every fragment pointer is `base + k*32`, so a 16 B aligned base put all of them on 16 B boundaries. Now `__align__(128)`. * Shared memory 34,816 -> 36,864 B, still under the 48 KB no-opt-in limit (asserted). Also renames N_SUBTILES -> N_SUB_TILES; the `typos` CI lane reads the concatenation as a misspelling of SUBTITLES. Found by verifying the WMMA contract against the shipped CUDA 12.4 headers and both specs rather than against memory. Nothing here has run. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01UMmjFy8TvsgypxVWNVhhC7 --- .../blockwise_fp8/blockwise_fp8_gemm_wmma.cu | 88 +++++++++++++------ 1 file changed, 61 insertions(+), 27 deletions(-) diff --git a/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu b/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu index 9291ad74f..347cc683e 100644 --- a/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu +++ b/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu @@ -152,12 +152,12 @@ constexpr int WMMA_K_DIM = 16; // 8 warps: 4 along M, 2 along N; each warp owns two 16-wide N sub-tiles. constexpr int WARPS_M = 4; constexpr int WARPS_N = 2; -constexpr int N_SUBTILES = 2; +constexpr int N_SUB_TILES = 2; constexpr int WARPS_PER_BLOCK = WARPS_M * WARPS_N; // 8 constexpr int BLOCK_THREADS = WARPS_PER_BLOCK * 32; // 256 constexpr int M_BLK = WARPS_M * WMMA_M_DIM; // 64 -constexpr int N_BLK = WARPS_N * N_SUBTILES * WMMA_N_DIM; // 64 +constexpr int N_BLK = WARPS_N * N_SUB_TILES * WMMA_N_DIM; // 64 // K_BLK is the scale-block length along K. The Rust dispatcher only selects // this kernel when `block_size_x % K_BLK == 0`, which is what makes the scale @@ -169,28 +169,56 @@ constexpr int WMMA_K_STEPS = K_BLK / WMMA_K_DIM; // 8 // Shared-memory row padding, in ELEMENTS of T. // -// 8, and it must stay a multiple of 8. `wmma::load_matrix_sync` requires the -// leading dimension to be a multiple of 16 bytes for 16-bit element types; -// (128 + 8) * 2 B = 272 B = 17 * 16 B, which satisfies it. This is the -// opposite of the right answer for the SCALAR kernel next door, where a +1 -// pad breaks a 4-way bank conflict -- there the loads are ordinary -// per-element `LDS.32` with no alignment rule, here they are fragment loads -// that have one. A +1 pad here would make the leading dimension 258 B, which -// is not a multiple of 16 B, and the load would be undefined. Do not -// "simplify" this to +1 to match the other kernel. -constexpr int SMEM_PAD = 8; -constexpr int A_LDM = K_BLK + SMEM_PAD; // 136 elements -constexpr int B_LDM = K_BLK + SMEM_PAD; // 136 elements +// 16, and it must stay a multiple of 16. READ THIS BEFORE CHANGING IT -- the +// two NVIDIA documents that govern it disagree, and the laxer one is the one +// people quote. +// +// The CUDA C++ Programming Guide 12.4 ยง7.24.1 says only: +// +// "mptr must be a 256-bit aligned pointer ... and [ldm] must be a +// multiple of 8 for __half element type or multiple of 4 for float +// element type. (i.e., multiple of 16 bytes in both cases)." +// +// The PTX ISA 12.4 ยง9.7.13.3.2 is STRICTER, and it is the binding one: +// +// "The starting address of each instance of the leading dimension (row +// or column) must be aligned with the size of the corresponding +// fragment in bytes." +// +// and works the example through for exactly our shape: for +// `wmma.load.a.sync.aligned.row.m16n16k16.f16` the fragment is 32 B (eight +// `.f16x2` elements), so "p is a multiple of 32" and "2*s is a multiple of +// 32" -- i.e. ldm must be a multiple of SIXTEEN __half elements, double what +// the C++ guide states. A +8 pad gives ldm = 136, and 136 * 2 B = 272 B is +// NOT a multiple of 32 B: it satisfies the guide and violates the ISA, which +// is the worst of both worlds because it compiles. +// +// +16 gives ldm = 144 elements = 288 B = 9 * 32 B, which satisfies both, for +// f16 and bf16 alike. Every fragment base pointer this kernel forms is then +// 32-B aligned as well (144 * 2 = 288 and 16 * 2 = 32 are both multiples of +// 32, and B_sh starts 18,432 B into a 16-B-aligned block). +// +// This is also the OPPOSITE of the right answer for the scalar kernel next +// door, where a +1 pad breaks a 4-way bank conflict. There the loads are +// ordinary per-element `LDS.32` with no alignment rule; here they are +// fragment loads that have one. Do not "harmonise" the two. +constexpr int SMEM_PAD = 16; +constexpr int A_LDM = K_BLK + SMEM_PAD; // 144 elements = 288 B = 9 * 32 B +constexpr int B_LDM = K_BLK + SMEM_PAD; // 144 elements +static_assert(A_LDM * 2 % 32 == 0, + "WMMA leading dimension must be a multiple of 32 bytes (PTX ISA " + "9.7.13.3.2), not merely of 16 as the C++ guide states"); +static_assert(B_LDM * 2 % 32 == 0, "see above"); // One K-tile of A and of B, plus (aliased over A) the f32 output staging tile. -constexpr int A_SMEM_ELEMS = M_BLK * A_LDM; // 8704 -constexpr int B_SMEM_ELEMS = N_BLK * B_LDM; // 8704 +constexpr int A_SMEM_ELEMS = M_BLK * A_LDM; // 64 * 144 = 9216 +constexpr int B_SMEM_ELEMS = N_BLK * B_LDM; // 64 * 144 = 9216 constexpr int C_SMEM_ELEMS = M_BLK * N_BLK; // 4096 floats // Bytes, for a 16-bit T. Static shared memory, so we stay under the 48 KB // per-block limit that needs no `cudaFuncSetAttribute` opt-in: -// (8704 + 8704) * 2 = 34,816 B. -// The f32 staging tile is 4096 * 4 = 16,384 B and ALIASES A_sh (17,408 B), +// (9216 + 9216) * 2 = 36,864 B. +// The f32 staging tile is 4096 * 4 = 16,384 B and ALIASES A_sh (18,432 B), // which is only safe because it is written after the K loop has finished with // A. The __syncthreads() before that reuse is load-bearing. constexpr int AB_SMEM_BYTES = (A_SMEM_ELEMS + B_SMEM_ELEMS) * 2; @@ -297,7 +325,13 @@ __launch_bounds__(BLOCK_THREADS) __global__ void fp8_matmul_wmma( "16x16x16 f32 accumulator fragment is not 8 registers/lane; " "the promotion loop below assumes it is"); - __shared__ __align__(16) char smem_raw[AB_SMEM_BYTES]; + // __align__(128), not 16. The PTX ISA rule quoted at SMEM_PAD requires the + // fragment base pointer itself to be a multiple of 32 B ("p is a multiple + // of 32"), and every pointer this kernel forms is `base + k*32`, so a + // 16-B-aligned base would put every one of them on a 16-B boundary and + // violate it. 128 covers that with room to spare and costs at most 127 B of + // a 36,864 B allocation. + __shared__ __align__(128) char smem_raw[AB_SMEM_BYTES]; T *A_sh = reinterpret_cast(smem_raw); T *B_sh = reinterpret_cast(smem_raw) + A_SMEM_ELEMS; @@ -317,10 +351,10 @@ __launch_bounds__(BLOCK_THREADS) __global__ void fp8_matmul_wmma( // Long-lived FP32 accumulators. `c_frag` accumulates ONE scale block and is // reset after every promotion; `acc` carries the scaled running total for // the whole K extent. - AccFrag c_frag[N_SUBTILES]; - float acc[N_SUBTILES][ACC_REGS]; + AccFrag c_frag[N_SUB_TILES]; + float acc[N_SUB_TILES][ACC_REGS]; #pragma unroll - for (int s = 0; s < N_SUBTILES; ++s) { + for (int s = 0; s < N_SUB_TILES; ++s) { fill_fragment(c_frag[s], 0.0f); #pragma unroll for (int i = 0; i < ACC_REGS; ++i) @@ -419,14 +453,14 @@ __launch_bounds__(BLOCK_THREADS) __global__ void fp8_matmul_wmma( A_LDM); #pragma unroll - for (int s = 0; s < N_SUBTILES; ++s) { + for (int s = 0; s < N_SUB_TILES; ++s) { // B_sh is [N][K]; as a K x N operand that is column-major with // leading dimension B_LDM. fragment b_frag; load_matrix_sync(b_frag, B_sh + - (warp_n_idx * N_SUBTILES + s) * WMMA_N_DIM * B_LDM + + (warp_n_idx * N_SUB_TILES + s) * WMMA_N_DIM * B_LDM + k_step * WMMA_K_DIM, B_LDM); mma_sync(c_frag[s], a_frag, b_frag, c_frag[s]); @@ -440,7 +474,7 @@ __launch_bounds__(BLOCK_THREADS) __global__ void fp8_matmul_wmma( { const float ws = __ldg(&weight_scale[scale_row_off + k_base / block_size_x]); #pragma unroll - for (int s = 0; s < N_SUBTILES; ++s) { + for (int s = 0; s < N_SUB_TILES; ++s) { #pragma unroll for (int i = 0; i < ACC_REGS; ++i) acc[s][i] += c_frag[s].x[i] * ws; @@ -460,12 +494,12 @@ __launch_bounds__(BLOCK_THREADS) __global__ void fp8_matmul_wmma( // writes C over it. float *C_sh = reinterpret_cast(smem_raw); #pragma unroll - for (int s = 0; s < N_SUBTILES; ++s) { + for (int s = 0; s < N_SUB_TILES; ++s) { #pragma unroll for (int i = 0; i < ACC_REGS; ++i) c_frag[s].x[i] = acc[s][i]; store_matrix_sync(C_sh + warp_m_idx * WMMA_M_DIM * N_BLK + - (warp_n_idx * N_SUBTILES + s) * WMMA_N_DIM, + (warp_n_idx * N_SUB_TILES + s) * WMMA_N_DIM, c_frag[s], N_BLK, mem_row_major); } __syncthreads(); From 197c847ba6d59f35e6d945dc80c2379073b9f5c1 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Fri, 21 Aug 2026 03:23:12 +0100 Subject: [PATCH 3/4] docs(ArcKernels): flag that lowering ARC_FP8_CUBLAS_MIN_M makes this kernel dead code The forward() dispatch chooses between the native FP8 path and dequantize_w() + cuBLASLt BEFORE reaching the WMMA kernel. Master's threshold is 512 so B=256 reaches it; PR #201 lowers it to 5, which would leave this kernel unreachable at every batch size that matters -- silently, since it still compiles and still passes its tests. The two are answers to the same question asked before and after the premise changed: cuBLASLt won the M=8..128 sweep because it had tensor cores and the native path did not. This kernel removes that asymmetry without paying the ~12.7 GB/step dequantize or the +8.48 GB of resident BF16 weights. Records the three-arm sweep that has to replace the two-arm one. Nothing has run; the expectation that (b) beats (a) is a derivation from bytes moved. UNVERIFIED ON HARDWARE -- never run. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01UMmjFy8TvsgypxVWNVhhC7 --- mistralrs-quant/src/blockwise_fp8/ops.rs | 38 ++++++++++++++++++++++++ 1 file changed, 38 insertions(+) diff --git a/mistralrs-quant/src/blockwise_fp8/ops.rs b/mistralrs-quant/src/blockwise_fp8/ops.rs index c75bfb249..2f8171f5c 100644 --- a/mistralrs-quant/src/blockwise_fp8/ops.rs +++ b/mistralrs-quant/src/blockwise_fp8/ops.rs @@ -944,6 +944,44 @@ fn fp8_gemv_max_m() -> usize { /// Set `ARC_NO_FP8_WMMA=1` to get back exactly the behaviour of the commit /// before this one. That is the control arm; run it first. /// +/// # โš ๏ธ THIS KERNEL IS UNREACHABLE IF `ARC_FP8_CUBLAS_MIN_M` IS LOWERED +/// +/// `BlockwiseFP8Linear::forward` decides between *this whole file* and the +/// dequantize + cuBLASLt fallback **before** it ever gets here: +/// +/// ```text +/// m_rows >= arc_fp8_cublas_min_m() -> dequantize_w() + cuBLASLt +/// otherwise -> the native FP8 path (this file) +/// ``` +/// +/// On master that constant is 512, so B = 256 reaches this kernel. A separate +/// change in flight (PR #201) lowers it to **5**, which routes everything +/// above the GEMV domain into the dequantize path and leaves this kernel +/// **dead code at every batch size that matters**. That is the exact +/// "wired-but-dead" failure this repo already tracks, and it would be silent: +/// the kernel compiles, the tests pass, and it never runs. +/// +/// The two changes are not actually in conflict โ€” they are answers to the +/// *same* question asked before and after the premise changed. The reason +/// cuBLASLt won the M = 8..128 sweep is that it had tensor cores and the +/// native FP8 path did not. This kernel removes that asymmetry, and it does +/// so while reading FP8 straight from HBM: no `dequantize_w()`, so none of +/// the ~12.7 GB/step of dequantize traffic and none of the +8.48 GB of +/// resident BF16 weights that the cuBLASLt arm pays for. +/// +/// So whoever merges these two must **re-run the threshold sweep with this +/// kernel present**, comparing three arms, not two: +/// +/// ```text +/// (a) dequantize_w() + cuBLASLt <- what PR #201 measured +/// (b) native FP8 tensor-core GEMM <- this file, ARC_NO_FP8_WMMA unset +/// (c) native scalar fp8_matmul_tiled <- ARC_NO_FP8_WMMA=1 +/// ``` +/// +/// Derivation, NOT a measurement: (b) should beat (a) because it moves +/// ~4.238 GB instead of ~12.7 GB and skips a full-model dequantize per step. +/// Nobody has run it. Do not merge a threshold on the two-arm result. +/// /// # Why there is no M threshold here /// /// There is an obvious temptation to gate this on some minimum M, since a From b0ef81dbb5e25d738334cf93ea112efc59667888 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Fri, 21 Aug 2026 03:25:12 +0100 Subject: [PATCH 4/4] docs(ArcKernels): name the occupancy limit to check first, before anyone measures At V4's B=256 shapes the grid is ~128 blocks against 132 SMs -- one wave, one block per SM, 8 warps of a possible 64. That is the most likely reason the kernel lands short of its 2.19 ms derived bound, and it is where a first profile should point. Lists the three knobs (smaller tiles, cp.async double buffering, split-K) and deliberately guesses none of them, for the same reason the dispatch threshold is left unset. UNVERIFIED ON HARDWARE -- never run. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01UMmjFy8TvsgypxVWNVhhC7 --- .../blockwise_fp8/blockwise_fp8_gemm_wmma.cu | 23 +++++++++++++++++++ 1 file changed, 23 insertions(+) diff --git a/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu b/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu index 347cc683e..3e518a373 100644 --- a/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu +++ b/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu @@ -93,6 +93,29 @@ * scales -- a numerics change and a separate kernel, deliberately not done * here. Named, not built. * + * WHERE TO LOOK FIRST IF IT UNDERPERFORMS THAT BOUND (unmeasured) + * --------------------------------------------------------------- + * Occupancy, most likely. At V4's B=256 decode shapes the grid is + * `ceil(N/64) x ceil(M/64)`; for M = 256 and N = 2048 that is 32 x 4 = 128 + * blocks against an H200's 132 SMs -- roughly ONE wave, one block per SM, + * 8 warps out of the 64 an SM can hold. There is very little left to hide + * global-load latency behind, and the tail effect is total: a second wave + * would cost as much as the first. + * + * The knobs, in the order worth trying, none of which have been tried: + * 1. Smaller M_BLK (32) or N_BLK (32) -- more blocks, more waves, better + * fill at the cost of re-reading operands. + * 2. More warps per block, or two K-tiles in flight (double-buffered A/B + * staging with `cp.async`) so the MMA pipe is not stalled on the load. + * 3. Split-K, which is the standard answer for a small-M, large-K GEMM and + * the regime decode actually lives in. + * + * These are deliberately NOT guessed at here. The tile shape is a set of + * compile-time constants chosen to match the working precedent in this tree, + * and picking a different one without a measurement would be the same mistake + * as the two frozen dispatch thresholds documented in blockwise_fp8/ops.rs. + * Sweep them on a box. + * * ALSO CHECKED, AND THE REASON THIS IS HAND-WRITTEN * ------------------------------------------------- * cuBLASLt gained exactly this layout -- `CUBLASLT_MATMUL_DESC_B_SCALE_MODE =