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..3e518a373 --- /dev/null +++ b/mistralrs-quant/kernels/blockwise_fp8/blockwise_fp8_gemm_wmma.cu @@ -0,0 +1,585 @@ +/** + * 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. + * + * 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 = + * 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_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_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 +// 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. +// +// 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; // 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: +// (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; +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"); + + // __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; + + 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_SUB_TILES]; + float acc[N_SUB_TILES][ACC_REGS]; +#pragma unroll + 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) + 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_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_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]); + } + } + + // ---- 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_SUB_TILES; ++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_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_SUB_TILES + 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..2f8171f5c 100644 --- a/mistralrs-quant/src/blockwise_fp8/ops.rs +++ b/mistralrs-quant/src/blockwise_fp8/ops.rs @@ -925,6 +925,133 @@ 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. +/// +/// # โš ๏ธ 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 +/// 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 +1139,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 +1202,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 +1273,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(