From 66e07326c3376c636db39cc13e6071e58f7be46f Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Tue, 18 Aug 2026 10:35:30 +0100 Subject: [PATCH 01/22] feat(ArcKV/Fp8): fused E4M3 quantize+dequantize CUDA kernel, default path The V4 B=1 decode step pays 43 blocking `cuMemcpyDtoHAsync_v2` per token (~109us each, 4.81 ms/token) inside `dsv4_kv_fp8::e4m3_codes_cpu`, which round-trips every layer's scaled K block through the HOST because candle has no CUDA F8E4M3 cast. Those copies are invisible to the obvious grep (`*Synchronize*` = 0.0 calls/step) and they make CUDA graph capture impossible: a graph cannot record a blocking D2H. The measured trap: the existing sync-free path (`ARC_GPU_ACT_QUANT=1`, `GpuApprox`) is SLOWER on H200 - interleaved A1 66.99 / B1 67.71 / A2 68.05 / B2 70.30 ms/token, with `kv_fp8_quant` 73.49 -> 134.18 us/call (+83%). It swaps one blocking copy for ~19 extra elementwise launches per layer, and this machine is bottlenecked on op count (9,131 launches/token, median kernel 1.18us, memory controller 4% utilized), not kernel speed. So `GpuApprox` is not the fix and is documented here as not being one. This adds `mistralrs_quant::arc_kvquant`: one fused kernel per direction, replacing ~11 candle ops on the quantize side and ~13 on the dequantize side with a single launch each, and removing the D2H entirely. The byte format is unchanged. Bit-parity with the CPU path is the bar, and it is obtained by construction rather than by hope: * the E4M3 rounding is a transcription of NVIDIA's `__nv_cvt_double_to_fp8(x, SATFINITE, E4M3)`, which is what the Rust `float8` crate ports and therefore what `F8E4M3::from_f32` computes; * `scale` reproduces `(amax / 448.0)?.affine(1.0, 1e-12)` exactly, including that candle lowers `Tensor / f64` to a MULTIPLY by the f32-rounded reciprocal; * mistralrs-quant compiles with `--use_fast_math`, so every float op is an explicit `__f*_rn` intrinsic (IEEE, unaffected by -prec-div/-ftz) and the amax reduction runs on absolute-value bit patterns as unsigned integers rather than through `fabsf`/`fmaxf`; * dequant indexes the SAME 256-entry `F8E4M3::from_bits(i).to_f32()` table the candle path fed to `index_select`. D33: the kernel ships a deliberate mutant (RNE replaced by truncation, everything else identical) reachable only from the parity test, so the comparison is shown to fail on a wrong kernel. The test also asserts the fused call counters moved - a parity check on this exact subsystem has already passed vacuously by comparing an implementation to itself. D14: the GPU parity test exits 2 (environment failure) rather than passing when no CUDA device is present. Co-Authored-By: Claude Opus 5 (1M context) --- mistralrs-core/src/models/dsv4_kv_fp8.rs | 243 ++++++++- .../kernels/arc_kvquant/arc_kvquant.cu | 461 ++++++++++++++++++ mistralrs-quant/src/arc_kvquant/ffi.rs | 59 +++ mistralrs-quant/src/arc_kvquant/mod.rs | 380 +++++++++++++++ mistralrs-quant/src/lib.rs | 1 + 5 files changed, 1139 insertions(+), 5 deletions(-) create mode 100644 mistralrs-quant/kernels/arc_kvquant/arc_kvquant.cu create mode 100644 mistralrs-quant/src/arc_kvquant/ffi.rs create mode 100644 mistralrs-quant/src/arc_kvquant/mod.rs diff --git a/mistralrs-core/src/models/dsv4_kv_fp8.rs b/mistralrs-core/src/models/dsv4_kv_fp8.rs index 5752bebf5..773e95861 100644 --- a/mistralrs-core/src/models/dsv4_kv_fp8.rs +++ b/mistralrs-core/src/models/dsv4_kv_fp8.rs @@ -48,6 +48,7 @@ use candle_core::{DType, Device, Result, Tensor, D}; use float8::F8E4M3; +use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{OnceLock, RwLock}; /// Block width of the activation quantizer, matching the reference's @@ -61,13 +62,31 @@ const E4M3_MAX: f64 = 448.0; /// Which arithmetic produces the E4M3 code. #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub(crate) enum KvQuantMode { + /// One fused device kernel per direction (`ArcQuant`, + /// `mistralrs_quant::arc_kvquant`). **The default.** Bit-identical to + /// [`Self::CpuExact`] — see that kernel's header for how each arithmetic + /// step is pinned to the candle op it replaces — but it is one launch + /// instead of ~11, and it removes the blocking device→host copy entirely. + /// + /// Falls back to [`Self::CpuExact`] when the tensor is not on CUDA, which + /// is what keeps the CPU unit tests below meaningful. + FusedDevice, /// Exact E4M3 via candle's CPU cast (candle has no CUDA `F8E4M3` cast — /// "named symbol not found"), at the price of one device sync per layer. + /// `ARC_KV_FP8_MODE=cpu`. CpuExact, /// On-device float arithmetic that reproduces E4M3's value grid with /// round-half-away-from-zero instead of round-half-to-even. Removes the /// sync; the FP8-trained model tolerates the sub-ULP-of-FP8 difference. /// `ARC_GPU_ACT_QUANT=1`. (RUN-161 throughput.) + /// + /// **Kept reachable, but it is NOT the fix and must not be shipped as one.** + /// Measured on H200 at B=1 it is *slower* than the sync it removes + /// (interleaved A1 66.99 / B1 67.71 / A2 68.05 / B2 70.30 ms/token, with + /// `kv_fp8_quant` 73.49 → 134.18 µs/call): it trades one blocking copy for + /// ~19 extra elementwise launches per layer, and this machine is bottlenecked + /// on op count, not on kernel speed. That measurement is the reason + /// [`Self::FusedDevice`] is one kernel rather than a chain of candle ops. GpuApprox, } @@ -79,15 +98,40 @@ impl KvQuantMode { pub(crate) fn from_env() -> Self { static MODE: OnceLock = OnceLock::new(); *MODE.get_or_init(|| { - if std::env::var_os("ARC_GPU_ACT_QUANT").is_some() { - Self::GpuApprox - } else { - Self::CpuExact + match std::env::var("ARC_KV_FP8_MODE") + .unwrap_or_default() + .to_ascii_lowercase() + .as_str() + { + "cpu" | "cpu_exact" => Self::CpuExact, + "gpu" | "gpu_approx" => Self::GpuApprox, + "fused" | "fused_device" => Self::FusedDevice, + _ if std::env::var_os("ARC_GPU_ACT_QUANT").is_some() => Self::GpuApprox, + _ => Self::FusedDevice, } }) } } +/// How many times the fused device kernels actually ran. +/// +/// D18/D33: a parity test that passes while these stay at zero has proved +/// nothing — it compared the candle path to itself. Every test that asserts the +/// fused path is correct also asserts these moved, and the fused path is the +/// only thing that touches them. +static FUSED_QUANT_CALLS: AtomicU64 = AtomicU64::new(0); +static FUSED_DEQUANT_CALLS: AtomicU64 = AtomicU64::new(0); + +/// Number of `quantize_k` calls served by the fused device kernel. +pub(crate) fn fused_quantize_calls() -> u64 { + FUSED_QUANT_CALLS.load(Ordering::Relaxed) +} + +/// Number of `V4PackedK::dequant` calls served by the fused device kernel. +pub(crate) fn fused_dequantize_calls() -> u64 { + FUSED_DEQUANT_CALLS.load(Ordering::Relaxed) +} + /// The 256 E4M3 values, indexed by code byte. Entry `i` is exactly what /// candle's `F8E4M3 -> F32` cast yields for the byte `i`, which is what makes /// [`V4PackedK::dequant`] bit-exact against the pre-packing implementation. @@ -165,7 +209,42 @@ impl V4PackedK { /// /// Bit-identical to what the pre-packing `act_quant_kv_nope` stored — see /// the module docs and `kv_fp8_roundtrip_is_bit_exact_vs_reference`. + /// + /// Served by the fused device kernel when it can be (default), and by the + /// candle op chain otherwise. The two produce the same bits, so the choice + /// is invisible to callers; it is exposed to tests through + /// [`Self::dequant_with`] so parity can be checked *between* them. pub(crate) fn dequant(&self, out_dtype: DType) -> Result { + self.dequant_with(KvQuantMode::from_env(), out_dtype) + } + + /// [`Self::dequant`] with the path chosen explicitly. + pub(crate) fn dequant_with(&self, mode: KvQuantMode, out_dtype: DType) -> Result { + if mode == KvQuantMode::FusedDevice + && out_dtype == self.side.dtype() + && mistralrs_quant::arc_kvquant::kv_fp8_fused_available(self.codes.device()) + { + let nope = self.codes.dim(D::Minus1)?; + let lut = e4m3_lut(self.codes.device())?; + let out = mistralrs_quant::arc_kvquant::kv_fp8_dequantize( + &self.codes, + &self.side, + &lut, + self.rope_dim, + KV_QUANT_BLOCK, + )?; + debug_assert_eq!(out.dim(D::Minus1)?, nope + self.rope_dim); + FUSED_DEQUANT_CALLS.fetch_add(1, Ordering::Relaxed); + return Ok(out); + } + self.dequant_candle(out_dtype) + } + + /// The candle op chain: ~13 launches, and the reference the fused kernel is + /// pinned against. Kept because it is the only implementation that runs on + /// CPU, and because a fused kernel with nothing to compare to is not a + /// verified kernel. + fn dequant_candle(&self, out_dtype: DType) -> Result { let (b, h, t, nope) = self.codes.dims4()?; let n_blocks = nope / KV_QUANT_BLOCK; let dev = self.codes.device(); @@ -222,6 +301,24 @@ pub(crate) fn quantize_k( } let n_blocks = nope / KV_QUANT_BLOCK; + // The whole point of this module's rewrite: one launch, no device→host + // round trip, and therefore nothing that blocks CUDA graph capture. Falls + // through to the candle chain on CPU (and on any dtype the kernel does not + // carry), which is what the parity tests compare against. + if mode == KvQuantMode::FusedDevice + && mistralrs_quant::arc_kvquant::kv_fp8_fused_available(k.device()) + && matches!(k.dtype(), DType::BF16 | DType::F16 | DType::F32) + { + let (codes, side) = + mistralrs_quant::arc_kvquant::kv_fp8_quantize(k, rope_dim, KV_QUANT_BLOCK)?; + FUSED_QUANT_CALLS.fetch_add(1, Ordering::Relaxed); + return Ok(Some(V4PackedK { + codes, + side, + rope_dim, + })); + } + let k_nope = k.narrow(D::Minus1, 0, nope)?; let k_rope = k.narrow(D::Minus1, nope, rope_dim)?; @@ -233,7 +330,10 @@ pub(crate) fn quantize_k( let scaled = kb.broadcast_div(&scale)?; let codes = match mode { - KvQuantMode::CpuExact => e4m3_codes_cpu(&scaled)?, + // `FusedDevice` only reaches here when the kernel could not serve the + // call (CPU tensor, or a dtype it does not carry). The exact path is + // the right fallback: it is what `FusedDevice` is bit-identical to. + KvQuantMode::FusedDevice | KvQuantMode::CpuExact => e4m3_codes_cpu(&scaled)?, KvQuantMode::GpuApprox => e4m3_codes_arith(&scaled)?, } .reshape((b, h, t, nope))?; @@ -490,6 +590,139 @@ mod tests { Ok(()) } + /// **The fused-kernel bit-parity proof, on hardware.** + /// + /// D14: this test is meaningless without a GPU, so it exits 2 (environment + /// failure) rather than passing when there is no CUDA device — a green run + /// with no device would be exactly the silent success this project has been + /// bitten by 13 times. + /// + /// D33: three independent quantities have to agree, and the mutant below + /// proves the comparison can fail. + /// + /// 1. `codes`/`side` bytes: fused kernel vs the candle chain + CPU cast. + /// 2. The reconstructed key tensor, computed **crosswise** — fused-quant + /// then candle-dequant against candle-quant then fused-dequant — so a + /// matching error would have to be present in both halves of two + /// different implementations. + /// 3. The full round trip against `reference_act_quant_kv_nope`, the + /// verbatim pre-packing implementation, run entirely on CPU. + /// + /// Plus the engagement assertion D18 demands: the fused call counters must + /// move, and must move only for the fused arm. A parity check on this exact + /// subsystem has already passed vacuously by comparing an implementation to + /// itself; only a launch counter caught it. + #[cfg(feature = "cuda")] + #[test] + fn kv_fp8_fused_is_bit_identical_to_cpu_exact() -> Result<()> { + const HEAD_DIM: usize = 512; + const ROPE: usize = 64; + + let dev = match Device::new_cuda(0) { + Ok(d) => d, + Err(e) => { + eprintln!("ENVFAIL: kv_fp8 fused parity needs a CUDA device: {e}"); + std::process::exit(2); + } + }; + + let k_cpu = fixture(2, 9, HEAD_DIM)?; + let k = k_cpu.to_device(&dev)?; + + let q0 = fused_quantize_calls(); + let packed_cpu = quantize_k(&k, ROPE, KvQuantMode::CpuExact)?.unwrap(); + assert_eq!( + fused_quantize_calls(), + q0, + "CpuExact must not reach the fused kernel — if it does, the two arms \ + are the same code and this test proves nothing" + ); + let packed_fused = quantize_k(&k, ROPE, KvQuantMode::FusedDevice)?.unwrap(); + assert_eq!( + fused_quantize_calls(), + q0 + 1, + "the fused arm did not reach the kernel: it silently fell back, so \ + any equality below is the candle path compared with itself" + ); + + // (1) The stored bytes. + let codes_cpu: Vec = packed_cpu.codes.flatten_all()?.to_vec1()?; + let codes_fused: Vec = packed_fused.codes.flatten_all()?.to_vec1()?; + assert_eq!( + codes_fused, codes_cpu, + "fused E4M3 codes differ from candle's CPU cast" + ); + assert_eq!( + bits(&packed_fused.side.to_device(&Device::Cpu)?)?, + bits(&packed_cpu.side.to_device(&Device::Cpu)?)?, + "fused `side` (rope tail + block amax) differs from the candle chain" + ); + + // (2) Crosswise reconstruction: each half of one implementation against + // the other half of the other. + let d0 = fused_dequantize_calls(); + let fused_then_candle = packed_fused.dequant_candle(DType::BF16)?; + assert_eq!( + fused_dequantize_calls(), + d0, + "dequant_candle must not reach the fused kernel" + ); + let candle_then_fused = packed_cpu.dequant_with(KvQuantMode::FusedDevice, DType::BF16)?; + assert_eq!( + fused_dequantize_calls(), + d0 + 1, + "the fused dequant silently fell back — nothing was tested" + ); + assert_eq!( + bits(&candle_then_fused.to_device(&Device::Cpu)?)?, + bits(&fused_then_candle.to_device(&Device::Cpu)?)?, + "fused dequant differs from the candle dequant chain" + ); + + // (3) Against the verbatim pre-packing implementation, on CPU. + let reference = reference_act_quant_kv_nope(&k_cpu, ROPE, false)?; + assert_eq!( + bits(&candle_then_fused.to_device(&Device::Cpu)?)?, + bits(&reference)?, + "the fused round trip does not reproduce what the pre-packing \ + implementation stored" + ); + + // D33: the comparison above must be able to fail. Same kernel, same + // scale, same layout, same reduction — only round-to-nearest-even + // replaced by truncation. + let (mutant_codes, mutant_side) = + mistralrs_quant::arc_kvquant::kv_fp8_quantize_mutant_for_test( + &k.contiguous()?, + ROPE, + KV_QUANT_BLOCK, + )?; + let mutant_codes: Vec = mutant_codes.flatten_all()?.to_vec1()?; + assert_eq!( + bits(&mutant_side.to_device(&Device::Cpu)?)?, + bits(&packed_cpu.side.to_device(&Device::Cpu)?)?, + "the mutant must differ ONLY in the rounding — if `side` differs too \ + it is not isolating the property under test" + ); + assert_ne!( + mutant_codes, codes_cpu, + "the negative control produced identical codes: this parity test \ + cannot fail and therefore proves nothing (D33)" + ); + let differing = mutant_codes + .iter() + .zip(codes_cpu.iter()) + .filter(|(a, b)| a != b) + .count(); + assert!( + differing * 20 > codes_cpu.len(), + "the negative control only perturbed {differing}/{} codes; a \ + comparison that weak could pass on luck", + codes_cpu.len() + ); + Ok(()) + } + /// The decode table must be the E4M3 grid itself, not an approximation of /// it: this is the single assumption `dequant`'s exactness rests on. #[test] diff --git a/mistralrs-quant/kernels/arc_kvquant/arc_kvquant.cu b/mistralrs-quant/kernels/arc_kvquant/arc_kvquant.cu new file mode 100644 index 000000000..7afb89f6e --- /dev/null +++ b/mistralrs-quant/kernels/arc_kvquant/arc_kvquant.cu @@ -0,0 +1,461 @@ +// Parent system: ArcQuant / TurboQuant — fused block-wise E4M3 activation +// quantizer for DeepSeek-V4's fused MQA key cache (ArcInfer / ArcKV / Fp8). +// +// WHY THIS FILE EXISTS +// ------------------- +// `dsv4_kv_fp8::e4m3_codes_cpu` quantized K by copying the scaled block to the +// HOST, casting there (candle has no CUDA F8E4M3 cast) and copying back. On the +// V4-Flash B=1 decode step that is 43 `cuMemcpyDtoHAsync_v2` per token at +// ~109 us of *host* time each (pageable staging => blocking) = 4.81 ms/token, +// and it makes CUDA graph capture impossible: a graph cannot record a blocking +// D2H. +// +// The obvious sync-free replacement (`KvQuantMode::GpuApprox`, ~19 extra +// elementwise candle ops per layer) was MEASURED SLOWER on H200 - in a machine +// bottlenecked on op count, one sync is cheaper than the launches it takes to +// avoid it. So the fix has to be *one* launch, not *fewer syncs*: these two +// kernels collapse the whole quantize (~11 candle ops) and dequantize (~13 +// candle ops) chains into a single kernel each. +// +// BIT-PARITY IS THE BAR - HOW IT IS OBTAINED +// ----------------------------------------- +// The stored bytes must be identical to what the CPU path already wrote, or +// every cache written before this change decodes differently. Each arithmetic +// step below therefore names the candle op it reproduces: +// +// amax = max |x| over the block <- abs() + max_keepdim() +// scale = amax * (float)(1/448) + 1e-12f <- `(amax / 448.0)?.affine(1,1e-12)` +// NOTE candle's `Tensor / f64` +// is `affine(1.0/rhs, 0.0)`, a +// MULTIPLY by the reciprocal, +// not a divide. +// code = E4M3(x / scale) <- broadcast_div + CPU cast +// value = lut[code] * scale <- index_select + broadcast_mul +// +// Two hazards are handled explicitly: +// +// 1. THIS CRATE COMPILES WITH `--use_fast_math` (mistralrs-quant/build.rs), so +// a bare `a / b` here would be `div.approx.f32` while candle-kernels (no +// fast math) emits IEEE `div.rn.f32`, and `-ftz=true` would flush denormals +// candle keeps. Every float operation below is therefore an explicit +// `__f*_rn` intrinsic, which the CUDA Math API defines as IEEE-754 and +// unaffected by `-prec-div`/`-ftz`/`-fmad`. `fabsf`/`fmaxf` are avoided in +// favour of integer ops on the bit pattern for the same reason. +// 2. The E4M3 rounding is a transcription of NVIDIA's +// `__nv_cvt_double_to_fp8(x, __NV_SATFINITE, __NV_E4M3)` (cuda_fp8.hpp), +// which is *also* what the Rust `float8` crate ports in `convert_to_fp8` +// and therefore what `F8E4M3::from_f32` - candle's CPU cast - computes. It +// is transcribed rather than #included so the parity claim is auditable +// against the Rust source in one place, and so it stays pure integer +// arithmetic that `--use_fast_math` cannot reach. +// +// The DEQUANT side does not convert code->float at all: it indexes the same +// 256-entry `f32` table the Rust side already builds with +// `F8E4M3::from_bits(i).to_f32()` and caches per device, so that half of the +// round trip is bit-exact by construction rather than by argument. + +#include +#include +#include + +namespace arc { + +// E4M3 max magnitude; `dsv4_kv_fp8::E4M3_MAX`. +#define ARC_KV_E4M3_MAX 448.0 +// The warp width the block reduction assumes. +#define ARC_KV_WARP 32 + +// --------------------------------------------------------------------------- +// Activation dtype <-> f32, matching candle's cast kernels exactly. +// f32 -> bf16 : candle `cast_f32_bf16` is `out[i] = inp[i]`, i.e. the +// `__nv_bfloat16(float)` constructor == `__float2bfloat16` +// (inline `cvt.rn.bf16.f32`, immune to -ftz). +// f32 -> f16 : likewise `__float2half` (inline `cvt.rn.f16.f32`). +// --------------------------------------------------------------------------- +template __device__ __forceinline__ float arc_to_f32(T v); +template <> __device__ __forceinline__ float arc_to_f32<__nv_bfloat16>(__nv_bfloat16 v) { + return __bfloat162float(v); +} +template <> __device__ __forceinline__ float arc_to_f32<__half>(__half v) { + return __half2float(v); +} +template <> __device__ __forceinline__ float arc_to_f32(float v) { return v; } + +template __device__ __forceinline__ T arc_from_f32(float v); +template <> __device__ __forceinline__ __nv_bfloat16 arc_from_f32<__nv_bfloat16>(float v) { + return __float2bfloat16(v); +} +template <> __device__ __forceinline__ __half arc_from_f32<__half>(float v) { + return __float2half(v); +} +template <> __device__ __forceinline__ float arc_from_f32(float v) { return v; } + +// --------------------------------------------------------------------------- +// scale = (amax / 448.0) then +1e-12, reproducing `dsv4_kv_fp8::block_scale`: +// (amax / E4M3_MAX)?.affine(1.0, 1e-12) +// candle lowers `Tensor / f64` to `affine(1.0/448.0, 0.0)` (bin_trait!(Div, .., +// |v| 1./v, ..)), i.e. a multiply by the f32-rounded reciprocal, and lowers +// `affine(mul, add)` to `x * mul + add` with `mul`/`add` first narrowed to the +// tensor dtype. `x * 1.0f` is exact, so the second affine is exactly an add, +// and FMA contraction of `x * mul + 0.0f` is exactly the multiply. Both are +// therefore reproduced by one `__fmul_rn` and one `__fadd_rn`. +// --------------------------------------------------------------------------- +__device__ __forceinline__ float arc_kv_block_scale(float amax) { + const float inv_max = (float)(1.0 / ARC_KV_E4M3_MAX); + const float eps = (float)1e-12; + return __fadd_rn(__fmul_rn(amax, inv_max), eps); +} + +// --------------------------------------------------------------------------- +// f32 -> E4M3 code byte. +// +// Transcription of NVIDIA `__nv_cvt_double_to_fp8(x, __NV_SATFINITE, +// __NV_E4M3)` (cuda_fp8.hpp), which the Rust `float8` crate ports verbatim as +// `convert_to_fp8`; `F8E4M3::from_f32(x)` is `convert_to_fp8(x as f64, +// SatFinite, E4M3)`, and `f32 -> f64` widening is exact. Every operation here +// is integer, so `--use_fast_math` cannot perturb it. +// +// Kept deliberately branch-shaped like the reference: the cost that mattered +// was a 109 us blocking memcpy, not a dozen integer ops. +// --------------------------------------------------------------------------- +__device__ __forceinline__ uint8_t arc_f32_to_e4m3_code(float xf) { + const double x = (double)xf; // exact + const uint64_t xbits = (uint64_t)__double_as_longlong(x); + + const uint8_t FP8_MAXNORM = 0x7EU; + const uint8_t FP8_MANTISSA_MASK = 0x07U; + const int FP8_EXP_BIAS = 7; + const int FP8_SIGNIFICAND_BITS = 4; + const uint64_t FP8_MINDENORM_O2 = 0x3F50000000000000ULL; + const uint64_t FP8_OVERFLOW_THRESHOLD = 0x407D000000000000ULL; + const uint64_t FP8_MINNORM = 0x3F90000000000000ULL; + const uint64_t DP_INF_BITS = 0x7FF0000000000000ULL; + const uint64_t FP8_DP_HALF_ULP = (uint64_t)1 << (53 - FP8_SIGNIFICAND_BITS - 1); + + const uint8_t sign = (uint8_t)((xbits >> 63) << 7); + const uint8_t exp = + (uint8_t)((int)((xbits >> 52) & 0x7FFULL) - 1023 + FP8_EXP_BIAS); + const uint8_t mantissa = + (uint8_t)((xbits >> (53 - FP8_SIGNIFICAND_BITS)) & (uint64_t)FP8_MANTISSA_MASK); + const uint64_t absx = xbits & 0x7FFFFFFFFFFFFFFFULL; + + uint8_t res; + if (absx <= FP8_MINDENORM_O2) { + // Zero or underflow. + res = 0U; + } else if (absx > DP_INF_BITS) { + // Preserve NaNs (E4M3 has a single NaN encoding). + res = 0x7FU; + } else if (absx > FP8_OVERFLOW_THRESHOLD) { + // SatFinite. + res = FP8_MAXNORM; + } else if (absx >= FP8_MINNORM) { + // Normal range, round-to-nearest-even. + res = (uint8_t)((uint8_t)(exp << (FP8_SIGNIFICAND_BITS - 1)) | mantissa); + const uint64_t round = xbits & ((FP8_DP_HALF_ULP << 1) - 1); + if ((round > FP8_DP_HALF_ULP) || + ((round == FP8_DP_HALF_ULP) && ((mantissa & 1U) != 0U))) { + res = (uint8_t)(res + 1U); + } + } else { + // Denormal range, round-to-nearest-even. + const uint8_t shift = (uint8_t)(1 - (int)exp); + const uint8_t man = (uint8_t)(mantissa | (uint8_t)(1U << (FP8_SIGNIFICAND_BITS - 1))); + res = (uint8_t)(man >> shift); + const uint64_t round = + (xbits | ((uint64_t)1 << (53 - 1))) & + ((FP8_DP_HALF_ULP << ((uint64_t)shift + 1)) - 1); + if ((round > (FP8_DP_HALF_ULP << (uint64_t)shift)) || + ((round == (FP8_DP_HALF_ULP << (uint64_t)shift)) && ((res & 1U) != 0U))) { + res = (uint8_t)(res + 1U); + } + } + return (uint8_t)(res | sign); +} + +// --------------------------------------------------------------------------- +// FUSED QUANTIZE. One launch replaces narrow/to_dtype/reshape/abs/max_keepdim/ +// affine/affine/broadcast_div/cast/reshape/cat. +// +// One CUDA block per token (a token is one (b, h, t) triple); one warp per +// 64-wide quant block. The amax reduction runs on the ABSOLUTE-VALUE BIT +// PATTERN as an unsigned integer: for non-negative IEEE floats the unsigned +// integer order is the float order, so `max` over `bits & 0x7fffffff` is +// exactly `max |x|` with no rounding, no `fabsf`, and no exposure to -ftz. +// --------------------------------------------------------------------------- +template +__global__ void arc_kv_fp8_quantize_kernel( + const T *__restrict__ k, // [ntok, head_dim] contiguous + uint8_t *__restrict__ codes, // [ntok, nope] + T *__restrict__ side, // [ntok, rope_dim + n_blocks] + const int head_dim, const int nope, const int rope_dim, const int n_blocks, + const int block_w, const long ntok) { + const int side_w = rope_dim + n_blocks; + const int lane = (int)(threadIdx.x & (ARC_KV_WARP - 1)); + const int warp = (int)(threadIdx.x / ARC_KV_WARP); + const int n_warps = (int)(blockDim.x / ARC_KV_WARP); + + for (long tok = (long)blockIdx.x; tok < ntok; tok += (long)gridDim.x) { + const T *krow = k + tok * (long)head_dim; + uint8_t *crow = codes + tok * (long)nope; + T *srow = side + tok * (long)side_w; + + for (int blk = warp; blk < n_blocks; blk += n_warps) { + const int base = blk * block_w; + + unsigned amax_bits = 0U; + for (int e = lane; e < block_w; e += ARC_KV_WARP) { + const unsigned b = __float_as_uint(arc_to_f32(krow[base + e])) & 0x7FFFFFFFU; + amax_bits = amax_bits > b ? amax_bits : b; + } +#pragma unroll + for (int off = ARC_KV_WARP / 2; off > 0; off >>= 1) { + const unsigned o = __shfl_xor_sync(0xFFFFFFFFU, amax_bits, off); + amax_bits = amax_bits > o ? amax_bits : o; + } + const float amax = __uint_as_float(amax_bits); + const float scale = arc_kv_block_scale(amax); + + for (int e = lane; e < block_w; e += ARC_KV_WARP) { + const float v = arc_to_f32(krow[base + e]); + crow[base + e] = arc_f32_to_e4m3_code(__fdiv_rn(v, scale)); + } + // `amax` IS one of the block's own elements, so narrowing it back to the + // activation dtype is exact - that is what lets dequant rebuild the + // identical scale. + if (lane == 0) { + srow[rope_dim + blk] = arc_from_f32(amax); + } + } + + // The RoPE'd tail is stored verbatim (RoPE is applied before the + // quantizer), and it is the FIRST `rope_dim` lanes of `side` - + // `cat(&[k_rope, amax_stored])`. + for (int i = (int)threadIdx.x; i < rope_dim; i += (int)blockDim.x) { + srow[i] = krow[nope + i]; + } + } +} + +// --------------------------------------------------------------------------- +// FUSED DEQUANTIZE. One launch replaces flatten/to_dtype/index_select/reshape/ +// narrow/to_dtype/reshape/affine/affine/broadcast_mul/reshape/to_dtype/narrow/ +// to_dtype/cat/contiguous. +// +// `lut` is the 256-entry f32 table the Rust side builds from +// `F8E4M3::from_bits(i).to_f32()` and caches per device - the SAME tensor the +// candle path passed to `index_select`, so the code->value half is bit-exact by +// construction. +// --------------------------------------------------------------------------- +template +__global__ void arc_kv_fp8_dequantize_kernel( + const uint8_t *__restrict__ codes, // [ntok, nope] + const T *__restrict__ side, // [ntok, rope_dim + n_blocks] + const float *__restrict__ lut, // [256] + T *__restrict__ out, // [ntok, head_dim] contiguous + const int head_dim, const int nope, const int rope_dim, const int n_blocks, + const int block_w, const long ntok) { + const int side_w = rope_dim + n_blocks; + const int lane = (int)(threadIdx.x & (ARC_KV_WARP - 1)); + const int warp = (int)(threadIdx.x / ARC_KV_WARP); + const int n_warps = (int)(blockDim.x / ARC_KV_WARP); + + for (long tok = (long)blockIdx.x; tok < ntok; tok += (long)gridDim.x) { + const uint8_t *crow = codes + tok * (long)nope; + const T *srow = side + tok * (long)side_w; + T *orow = out + tok * (long)head_dim; + + for (int blk = warp; blk < n_blocks; blk += n_warps) { + const int base = blk * block_w; + const float amax = arc_to_f32(srow[rope_dim + blk]); + const float scale = arc_kv_block_scale(amax); + for (int e = lane; e < block_w; e += ARC_KV_WARP) { + orow[base + e] = arc_from_f32(__fmul_rn(lut[crow[base + e]], scale)); + } + } + + // `cat(&[k_nope, k_rope])`: the RoPE'd tail follows the nope dims. + for (int i = (int)threadIdx.x; i < rope_dim; i += (int)blockDim.x) { + orow[nope + i] = srow[i]; + } + } +} + +} // namespace arc + +// --------------------------------------------------------------------------- +// C ABI. `stream_ptr` is candle's stream so the launch is recorded into a CUDA +// graph captured on it (cf. kvwrite.cu; rotary.cu's hardcoded stream 0 is NOT +// capturable). Grid is capped so a long prefill loops instead of launching an +// unbounded grid. +// --------------------------------------------------------------------------- +#define ARC_KV_MAX_BLOCKS 65535 + +#define ARC_KV_LAUNCH_QUANT(T) \ + arc::arc_kv_fp8_quantize_kernel<<>>( \ + reinterpret_cast(k), codes, reinterpret_cast(side), \ + head_dim, nope, rope_dim, n_blocks, block_w, ntok) + +#define ARC_KV_LAUNCH_DEQUANT(T) \ + arc::arc_kv_fp8_dequantize_kernel<<>>( \ + codes, reinterpret_cast(side), lut, \ + reinterpret_cast(out), head_dim, nope, rope_dim, n_blocks, block_w, \ + ntok) + +extern "C" void arc_kv_fp8_quantize( + const void *k, // [ntok, head_dim] activation dtype + uint8_t *codes, // [ntok, nope] + void *side, // [ntok, rope_dim + n_blocks] + int32_t head_dim, int32_t nope, int32_t rope_dim, int32_t n_blocks, + int32_t block_w, int64_t ntok, + void *stream_ptr, // cudaStream_t; null => default stream + uint32_t dtype // 0 => f16, 1 => bf16, 2 => f32 +) { + if (ntok <= 0 || n_blocks <= 0) { + return; + } + // One warp per quant block, capped at 256 threads (8 warps). + int warps = n_blocks < 8 ? n_blocks : 8; + const int threads = warps * ARC_KV_WARP; + long blocks = ntok < (long)ARC_KV_MAX_BLOCKS ? ntok : (long)ARC_KV_MAX_BLOCKS; + dim3 grid((unsigned)blocks); + dim3 block((unsigned)threads); + const cudaStream_t stream = reinterpret_cast(stream_ptr); + + if (dtype == 0) { + ARC_KV_LAUNCH_QUANT(__half); + } else if (dtype == 1) { + ARC_KV_LAUNCH_QUANT(__nv_bfloat16); + } else if (dtype == 2) { + ARC_KV_LAUNCH_QUANT(float); + } +} + +extern "C" void arc_kv_fp8_dequantize( + const uint8_t *codes, // [ntok, nope] + const void *side, // [ntok, rope_dim + n_blocks] + const float *lut, // [256] + void *out, // [ntok, head_dim] + int32_t head_dim, int32_t nope, int32_t rope_dim, int32_t n_blocks, + int32_t block_w, int64_t ntok, void *stream_ptr, uint32_t dtype) { + if (ntok <= 0 || n_blocks <= 0) { + return; + } + int warps = n_blocks < 8 ? n_blocks : 8; + const int threads = warps * ARC_KV_WARP; + long blocks = ntok < (long)ARC_KV_MAX_BLOCKS ? ntok : (long)ARC_KV_MAX_BLOCKS; + dim3 grid((unsigned)blocks); + dim3 block((unsigned)threads); + const cudaStream_t stream = reinterpret_cast(stream_ptr); + + if (dtype == 0) { + ARC_KV_LAUNCH_DEQUANT(__half); + } else if (dtype == 1) { + ARC_KV_LAUNCH_DEQUANT(__nv_bfloat16); + } else if (dtype == 2) { + ARC_KV_LAUNCH_DEQUANT(float); + } +} + +// --------------------------------------------------------------------------- +// D33 NEGATIVE CONTROL. Same as `arc_kv_fp8_quantize` except the rounding is +// truncated (round-toward-zero) instead of round-to-nearest-even. It exists so +// the bit-parity test can be shown to FAIL on a wrong kernel: a parity check on +// this exact subsystem has already passed vacuously by comparing an +// implementation to itself. +// +// Nothing in the serving path calls this - it is reachable only from +// `arc_kvquant::mutant_quantize_for_parity_test`. +// --------------------------------------------------------------------------- +namespace arc { + +// Identical to `arc_f32_to_e4m3_code` except that in the normal range it +// TRUNCATES the mantissa instead of rounding to nearest even - the single most +// plausible way to get an E4M3 quantizer subtly wrong. Everything else (the +// scale, the amax reduction, the layout, the rope copy) is untouched, so a +// parity test that survives this mutant is testing nothing. +__device__ __forceinline__ uint8_t arc_f32_to_e4m3_code_truncating(float xf) { + const double x = (double)xf; + const uint64_t xbits = (uint64_t)__double_as_longlong(x); + const uint64_t absx = xbits & 0x7FFFFFFFFFFFFFFFULL; + const uint64_t FP8_MINNORM = 0x3F90000000000000ULL; + const uint64_t FP8_OVERFLOW_THRESHOLD = 0x407D000000000000ULL; + if (absx < FP8_MINNORM || absx > FP8_OVERFLOW_THRESHOLD) { + return arc_f32_to_e4m3_code(xf); + } + const uint8_t exp = (uint8_t)((int)((xbits >> 52) & 0x7FFULL) - 1023 + 7); + const uint8_t mantissa = (uint8_t)((xbits >> 49) & 0x7ULL); + const uint8_t sign = (uint8_t)((xbits >> 63) << 7); + return (uint8_t)((uint8_t)((uint8_t)(exp << 3) | mantissa) | sign); +} + +template +__global__ void arc_kv_fp8_quantize_mutant_kernel( + const T *__restrict__ k, uint8_t *__restrict__ codes, T *__restrict__ side, + const int head_dim, const int nope, const int rope_dim, const int n_blocks, + const int block_w, const long ntok) { + const int side_w = rope_dim + n_blocks; + const int lane = (int)(threadIdx.x & (ARC_KV_WARP - 1)); + const int warp = (int)(threadIdx.x / ARC_KV_WARP); + const int n_warps = (int)(blockDim.x / ARC_KV_WARP); + + for (long tok = (long)blockIdx.x; tok < ntok; tok += (long)gridDim.x) { + const T *krow = k + tok * (long)head_dim; + uint8_t *crow = codes + tok * (long)nope; + T *srow = side + tok * (long)side_w; + for (int blk = warp; blk < n_blocks; blk += n_warps) { + const int base = blk * block_w; + unsigned amax_bits = 0U; + for (int e = lane; e < block_w; e += ARC_KV_WARP) { + const unsigned b = __float_as_uint(arc_to_f32(krow[base + e])) & 0x7FFFFFFFU; + amax_bits = amax_bits > b ? amax_bits : b; + } +#pragma unroll + for (int off = ARC_KV_WARP / 2; off > 0; off >>= 1) { + const unsigned o = __shfl_xor_sync(0xFFFFFFFFU, amax_bits, off); + amax_bits = amax_bits > o ? amax_bits : o; + } + const float amax = __uint_as_float(amax_bits); + const float scale = arc_kv_block_scale(amax); + for (int e = lane; e < block_w; e += ARC_KV_WARP) { + const float v = arc_to_f32(krow[base + e]); + crow[base + e] = arc_f32_to_e4m3_code_truncating(__fdiv_rn(v, scale)); + } + if (lane == 0) { + srow[rope_dim + blk] = arc_from_f32(amax); + } + } + for (int i = (int)threadIdx.x; i < rope_dim; i += (int)blockDim.x) { + srow[i] = krow[nope + i]; + } + } +} + +} // namespace arc + +#define ARC_KV_LAUNCH_MUTANT(T) \ + arc::arc_kv_fp8_quantize_mutant_kernel<<>>( \ + reinterpret_cast(k), codes, reinterpret_cast(side), \ + head_dim, nope, rope_dim, n_blocks, block_w, ntok) + +extern "C" void arc_kv_fp8_quantize_mutant( + const void *k, uint8_t *codes, void *side, int32_t head_dim, int32_t nope, + int32_t rope_dim, int32_t n_blocks, int32_t block_w, int64_t ntok, + void *stream_ptr, uint32_t dtype) { + if (ntok <= 0 || n_blocks <= 0) { + return; + } + int warps = n_blocks < 8 ? n_blocks : 8; + const int threads = warps * ARC_KV_WARP; + long blocks = ntok < (long)ARC_KV_MAX_BLOCKS ? ntok : (long)ARC_KV_MAX_BLOCKS; + dim3 grid((unsigned)blocks); + dim3 block((unsigned)threads); + const cudaStream_t stream = reinterpret_cast(stream_ptr); + if (dtype == 0) { + ARC_KV_LAUNCH_MUTANT(__half); + } else if (dtype == 1) { + ARC_KV_LAUNCH_MUTANT(__nv_bfloat16); + } else if (dtype == 2) { + ARC_KV_LAUNCH_MUTANT(float); + } +} diff --git a/mistralrs-quant/src/arc_kvquant/ffi.rs b/mistralrs-quant/src/arc_kvquant/ffi.rs new file mode 100644 index 000000000..dafc5bd42 --- /dev/null +++ b/mistralrs-quant/src/arc_kvquant/ffi.rs @@ -0,0 +1,59 @@ +use core::ffi::c_void; + +extern "C" { + /// Fused block-wise E4M3 quantize. See `kernels/arc_kvquant/arc_kvquant.cu`. + /// + /// * `k` - `[ntok, head_dim]` contiguous, activation dtype. + /// * `codes` - `[ntok, nope]` U8 out. + /// * `side` - `[ntok, rope_dim + n_blocks]` activation dtype out + /// (`cat(&[k_rope, amax])`). + /// * `stream` - candle's stream, so the launch is recordable into a CUDA + /// graph captured on it. + /// * `dtype` - 0 => f16, 1 => bf16, 2 => f32. + pub(crate) fn arc_kv_fp8_quantize( + k: *const c_void, + codes: *mut u8, + side: *mut c_void, + head_dim: i32, + nope: i32, + rope_dim: i32, + n_blocks: i32, + block_w: i32, + ntok: i64, + stream: *mut c_void, + dtype: u32, + ); + + /// Fused dequantize back to `[ntok, head_dim]` (`cat(&[k_nope, k_rope])`). + /// `lut` is the 256-entry F32 `F8E4M3::from_bits(i).to_f32()` table. + pub(crate) fn arc_kv_fp8_dequantize( + codes: *const u8, + side: *const c_void, + lut: *const f32, + out: *mut c_void, + head_dim: i32, + nope: i32, + rope_dim: i32, + n_blocks: i32, + block_w: i32, + ntok: i64, + stream: *mut c_void, + dtype: u32, + ); + + /// D33 negative control: `arc_kv_fp8_quantize` with round-to-nearest-even + /// replaced by truncation. Test-only; nothing in the serving path calls it. + pub(crate) fn arc_kv_fp8_quantize_mutant( + k: *const c_void, + codes: *mut u8, + side: *mut c_void, + head_dim: i32, + nope: i32, + rope_dim: i32, + n_blocks: i32, + block_w: i32, + ntok: i64, + stream: *mut c_void, + dtype: u32, + ); +} diff --git a/mistralrs-quant/src/arc_kvquant/mod.rs b/mistralrs-quant/src/arc_kvquant/mod.rs new file mode 100644 index 000000000..887043b68 --- /dev/null +++ b/mistralrs-quant/src/arc_kvquant/mod.rs @@ -0,0 +1,380 @@ +//! Parent system: ArcQuant / TurboQuant — fused block-wise E4M3 activation +//! quantizer for the DeepSeek-V4 fused MQA key cache (ArcInfer / ArcKV / Fp8). +//! +//! One launch each for quantize and dequantize, replacing the ~11 and ~13 +//! candle ops the caller used to issue — and, critically, replacing the +//! device→host→device round trip that `dsv4_kv_fp8::e4m3_codes_cpu` performs +//! because candle has no CUDA `F8E4M3` cast. That round trip is 43 blocking +//! `cuMemcpyDtoHAsync_v2` per V4 decode token and it makes CUDA graph capture +//! impossible: a graph cannot record a blocking D2H. +//! +//! The byte format is unchanged — same `codes`/`side` layout, same bits. See +//! `kernels/arc_kvquant/arc_kvquant.cu` for how bit-parity with the CPU path is +//! obtained, and for the deliberate mutant used to prove the parity test can +//! fail. + +#[cfg(feature = "cuda")] +mod ffi; + +#[cfg(feature = "cuda")] +mod cuda_impl { + use candle_core::{ + cuda_backend::{cudarc::driver::DeviceRepr, CudaDType}, + CudaDevice, CudaStorage, DType, Device, Result, Shape, Storage, Tensor, + }; + use core::ffi::c_void; + use half::{bf16, f16}; + + use crate::utils::slice_ptr; + + /// `dtype` discriminant shared with the C ABI. + fn dtype_id(dt: DType) -> Result { + match dt { + DType::F16 => Ok(0), + DType::BF16 => Ok(1), + DType::F32 => Ok(2), + other => candle_core::bail!("arc-kvquant: unsupported activation dtype {other:?}"), + } + } + + /// Geometry shared by both kernels, validated once. + struct Geom { + b: usize, + h: usize, + t: usize, + head_dim: usize, + nope: usize, + rope_dim: usize, + n_blocks: usize, + block_w: usize, + } + + impl Geom { + fn ntok(&self) -> usize { + self.b * self.h * self.t + } + fn side_w(&self) -> usize { + self.rope_dim + self.n_blocks + } + } + + fn cuda_device(t: &Tensor, what: &str) -> Result { + match t.device() { + Device::Cuda(d) => Ok(d.clone()), + _ => candle_core::bail!("arc-kvquant: {what} must live on CUDA"), + } + } + + /// Whether the fused kernels can serve `dev`. Callers fall back to the + /// candle op chain when this is false rather than failing. + pub fn kv_fp8_fused_available(dev: &Device) -> bool { + matches!(dev, Device::Cuda(_)) + } + + fn geometry(k: &Tensor, rope_dim: usize, block_w: usize) -> Result { + let (b, h, t, head_dim) = k.dims4()?; + if rope_dim > head_dim { + candle_core::bail!("arc-kvquant: rope_dim {rope_dim} > head_dim {head_dim}"); + } + let nope = head_dim - rope_dim; + if block_w == 0 || nope == 0 || nope % block_w != 0 { + candle_core::bail!( + "arc-kvquant: nope {nope} is not a whole number of {block_w}-wide blocks" + ); + } + Ok(Geom { + b, + h, + t, + head_dim, + nope, + rope_dim, + n_blocks: nope / block_w, + block_w, + }) + } + + /// Fused block-wise E4M3 quantize of `k`'s non-RoPE dims. + /// + /// `k` is `[B, H, T, head_dim]` in the activation dtype. Returns + /// `(codes, side)`: + /// + /// * `codes` — `[B, H, T, head_dim - rope_dim]` U8, one E4M3 code per dim; + /// * `side` — `[B, H, T, rope_dim + n_blocks]` in `k`'s dtype: the RoPE'd + /// tail verbatim followed by each 64-wide block's `amax`. + /// + /// Bit-identical to `dsv4_kv_fp8::quantize_k` under `KvQuantMode::CpuExact`. + /// That is the contract; `kv_fp8_fused_is_bit_identical_to_cpu_exact` pins + /// it on hardware, and the mutant below proves that test can fail. + pub fn kv_fp8_quantize( + k: &Tensor, + rope_dim: usize, + block_w: usize, + ) -> Result<(Tensor, Tensor)> { + quantize_inner(k, rope_dim, block_w, false) + } + + /// D33 negative control: identical to [`kv_fp8_quantize`] except the E4M3 + /// rounding truncates instead of rounding to nearest even. Exists purely so + /// the parity test can be shown to fail. Nothing in the serving path calls + /// it. + pub fn kv_fp8_quantize_mutant_for_test( + k: &Tensor, + rope_dim: usize, + block_w: usize, + ) -> Result<(Tensor, Tensor)> { + quantize_inner(k, rope_dim, block_w, true) + } + + fn quantize_inner( + k: &Tensor, + rope_dim: usize, + block_w: usize, + mutant: bool, + ) -> Result<(Tensor, Tensor)> { + let g = geometry(k, rope_dim, block_w)?; + let dev = cuda_device(k, "k")?; + let did = dtype_id(k.dtype())?; + let k = k.contiguous()?; + match k.dtype() { + DType::F16 => quantize_t::(&k, &g, &dev, did, mutant), + DType::BF16 => quantize_t::(&k, &g, &dev, did, mutant), + DType::F32 => quantize_t::(&k, &g, &dev, did, mutant), + other => candle_core::bail!("arc-kvquant: unsupported activation dtype {other:?}"), + } + } + + fn quantize_t( + k: &Tensor, + g: &Geom, + dev: &CudaDevice, + did: u32, + mutant: bool, + ) -> Result<(Tensor, Tensor)> { + let ntok = g.ntok(); + // `alloc` rather than `alloc_zeros`: both outputs are fully written, and + // the memset would be one more device op in the exact place where op + // count is the disease. + let codes_buf = unsafe { dev.alloc::(ntok * g.nope)? }; + let side_buf = unsafe { dev.alloc::(ntok * g.side_w())? }; + + let (k_storage, k_layout) = k.storage_and_layout(); + let k_s = match &*k_storage { + Storage::Cuda(s) => s, + _ => candle_core::bail!("arc-kvquant: k must be CUDA storage"), + }; + let (k_ptr, _k_guard) = slice_ptr(k_s.as_cuda_slice::()?, k_layout.start_offset()); + let (codes_ptr, _codes_guard) = slice_ptr(&codes_buf, 0); + let (side_ptr, _side_guard) = slice_ptr(&side_buf, 0); + + let launch = if mutant { + super::ffi::arc_kv_fp8_quantize_mutant + } else { + super::ffi::arc_kv_fp8_quantize + }; + unsafe { + launch( + k_ptr as *const c_void, + codes_ptr as *mut u8, + side_ptr as *mut c_void, + g.head_dim as i32, + g.nope as i32, + g.rope_dim as i32, + g.n_blocks as i32, + g.block_w as i32, + ntok as i64, + dev.cuda_stream().cu_stream() as *mut c_void, + did, + ) + }; + + drop(_k_guard); + drop(_codes_guard); + drop(_side_guard); + drop(k_storage); + + let codes = Tensor::from(( + Storage::Cuda(CudaStorage::wrap_cuda_slice(codes_buf, dev.clone())), + Shape::from_dims(&[g.b, g.h, g.t, g.nope]), + )); + let side = Tensor::from(( + Storage::Cuda(CudaStorage::wrap_cuda_slice(side_buf, dev.clone())), + Shape::from_dims(&[g.b, g.h, g.t, g.side_w()]), + )); + Ok((codes, side)) + } + + /// Fused dequantize: `codes` + `side` back to `[B, H, T, head_dim]` in + /// `side`'s dtype. + /// + /// `lut` is the 256-entry F32 table built from + /// `F8E4M3::from_bits(i).to_f32()` — the *same* tensor the candle path fed + /// to `index_select`, which is what makes the code→value half of the round + /// trip bit-exact by construction rather than by argument. + pub fn kv_fp8_dequantize( + codes: &Tensor, + side: &Tensor, + lut: &Tensor, + rope_dim: usize, + block_w: usize, + ) -> Result { + let (b, h, t, nope) = codes.dims4()?; + let (sb, sh, st, side_w) = side.dims4()?; + if (sb, sh, st) != (b, h, t) { + candle_core::bail!( + "arc-kvquant: side dims {:?} do not match codes {:?}", + side.dims(), + codes.dims() + ); + } + if block_w == 0 || nope == 0 || nope % block_w != 0 { + candle_core::bail!("arc-kvquant: nope {nope} not a multiple of block_w {block_w}"); + } + let n_blocks = nope / block_w; + if side_w != rope_dim + n_blocks { + candle_core::bail!( + "arc-kvquant: side width {side_w} != rope_dim {rope_dim} + n_blocks {n_blocks}" + ); + } + if codes.dtype() != DType::U8 { + candle_core::bail!("arc-kvquant: codes must be U8, got {:?}", codes.dtype()); + } + if lut.dtype() != DType::F32 || lut.elem_count() != 256 { + candle_core::bail!("arc-kvquant: lut must be 256 F32 entries"); + } + let g = Geom { + b, + h, + t, + head_dim: nope + rope_dim, + nope, + rope_dim, + n_blocks, + block_w, + }; + let dev = cuda_device(side, "side")?; + let did = dtype_id(side.dtype())?; + // A `narrow` on the sequence dim leaves a contiguous view (with a start + // offset) whenever B*H == 1 — the decode case this exists for — so + // `contiguous()` is free there. For B*H > 1 it materialises the window, + // which is what threading 4-D strides into the kernel would avoid; that + // is a batch-path follow-up, not a correctness gap. + let codes = codes.contiguous()?; + let side = side.contiguous()?; + let lut = lut.contiguous()?; + match side.dtype() { + DType::F16 => dequantize_t::(&codes, &side, &lut, &g, &dev, did), + DType::BF16 => dequantize_t::(&codes, &side, &lut, &g, &dev, did), + DType::F32 => dequantize_t::(&codes, &side, &lut, &g, &dev, did), + other => candle_core::bail!("arc-kvquant: unsupported activation dtype {other:?}"), + } + } + + fn dequantize_t( + codes: &Tensor, + side: &Tensor, + lut: &Tensor, + g: &Geom, + dev: &CudaDevice, + did: u32, + ) -> Result { + let ntok = g.ntok(); + let out_buf = unsafe { dev.alloc::(ntok * g.head_dim)? }; + + let (codes_storage, codes_layout) = codes.storage_and_layout(); + let codes_s = match &*codes_storage { + Storage::Cuda(s) => s, + _ => candle_core::bail!("arc-kvquant: codes must be CUDA storage"), + }; + let (side_storage, side_layout) = side.storage_and_layout(); + let side_s = match &*side_storage { + Storage::Cuda(s) => s, + _ => candle_core::bail!("arc-kvquant: side must be CUDA storage"), + }; + let (lut_storage, lut_layout) = lut.storage_and_layout(); + let lut_s = match &*lut_storage { + Storage::Cuda(s) => s, + _ => candle_core::bail!("arc-kvquant: lut must be CUDA storage"), + }; + + let (codes_ptr, _codes_guard) = + slice_ptr(codes_s.as_cuda_slice::()?, codes_layout.start_offset()); + let (side_ptr, _side_guard) = + slice_ptr(side_s.as_cuda_slice::()?, side_layout.start_offset()); + let (lut_ptr, _lut_guard) = + slice_ptr(lut_s.as_cuda_slice::()?, lut_layout.start_offset()); + let (out_ptr, _out_guard) = slice_ptr(&out_buf, 0); + + unsafe { + super::ffi::arc_kv_fp8_dequantize( + codes_ptr as *const u8, + side_ptr as *const c_void, + lut_ptr as *const f32, + out_ptr as *mut c_void, + g.head_dim as i32, + g.nope as i32, + g.rope_dim as i32, + g.n_blocks as i32, + g.block_w as i32, + ntok as i64, + dev.cuda_stream().cu_stream() as *mut c_void, + did, + ) + }; + + drop(_codes_guard); + drop(_side_guard); + drop(_lut_guard); + drop(_out_guard); + drop(codes_storage); + drop(side_storage); + drop(lut_storage); + + Ok(Tensor::from(( + Storage::Cuda(CudaStorage::wrap_cuda_slice(out_buf, dev.clone())), + Shape::from_dims(&[g.b, g.h, g.t, g.head_dim]), + ))) + } +} + +#[cfg(feature = "cuda")] +pub use cuda_impl::*; + +#[cfg(not(feature = "cuda"))] +mod stub { + use candle_core::{Device, Result, Tensor}; + + pub fn kv_fp8_quantize( + _k: &Tensor, + _rope_dim: usize, + _block_w: usize, + ) -> Result<(Tensor, Tensor)> { + candle_core::bail!("arc-kvquant: fused FP8 KV quantize requires the `cuda` feature") + } + + pub fn kv_fp8_quantize_mutant_for_test( + _k: &Tensor, + _rope_dim: usize, + _block_w: usize, + ) -> Result<(Tensor, Tensor)> { + candle_core::bail!("arc-kvquant: fused FP8 KV quantize requires the `cuda` feature") + } + + pub fn kv_fp8_dequantize( + _codes: &Tensor, + _side: &Tensor, + _lut: &Tensor, + _rope_dim: usize, + _block_w: usize, + ) -> Result { + candle_core::bail!("arc-kvquant: fused FP8 KV dequantize requires the `cuda` feature") + } + + /// Never available without the `cuda` feature. + pub fn kv_fp8_fused_available(_dev: &Device) -> bool { + false + } +} + +#[cfg(not(feature = "cuda"))] +pub use stub::*; diff --git a/mistralrs-quant/src/lib.rs b/mistralrs-quant/src/lib.rs index bbe893b15..2bb42c3bd 100644 --- a/mistralrs-quant/src/lib.rs +++ b/mistralrs-quant/src/lib.rs @@ -16,6 +16,7 @@ use pertensor_fp8::pertensor_fp8_linear_b; mod metal_kernels; mod afq; +pub mod arc_kvquant; mod bitsandbytes; mod blockwise_fp8; pub mod calibration; From e10cdea8c149b4fd20a51ab0d9aea903a3c43533 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Tue, 18 Aug 2026 10:40:10 +0100 Subject: [PATCH 02/22] fix(ArcKV/Fp8): inline-PTX IEEE ops (fast-math rewrote __f*_rn to .ftz) + box paths The emitted PTX showed nvcc 13.1 turning __fmul_rn/__fadd_rn/__fdiv_rn into mul.rn.ftz.f32 / div.rn.ftz.f32 / add.rn.ftz.f32 under --use_fast_math, which candle-kernels (no fast math) does not do. Replaced with inline PTX, which no optimisation flag rewrites, so 'grep -c .ftz.f32' over the PTX is the audit. Co-Authored-By: Claude Opus 5 (1M context) --- arc-tools/measure_kv_fp8_fused.sh | 223 ++++++++++++++++++ .../kernels/arc_kvquant/arc_kvquant.cu | 55 ++++- 2 files changed, 268 insertions(+), 10 deletions(-) create mode 100755 arc-tools/measure_kv_fp8_fused.sh diff --git a/arc-tools/measure_kv_fp8_fused.sh b/arc-tools/measure_kv_fp8_fused.sh new file mode 100755 index 000000000..fcf03815a --- /dev/null +++ b/arc-tools/measure_kv_fp8_fused.sh @@ -0,0 +1,223 @@ +#!/usr/bin/env bash +# ArcKV/Fp8 — measure the fused E4M3 quantize+dequantize kernel against the +# CPU round trip it replaces. Runs ON THE BOX. Nothing here is a proxy: every +# number is produced by the real binary on the real model. +# +# WHAT IT PRODUCES (the four numbers the change is judged on) +# 1. `cuMemcpyDtoHAsync_v2` calls per decode step, before and after. +# The 44/step are the thing being removed and the reason CUDA graph +# capture is impossible. NOTE `*Synchronize*` is 0.0 calls/step in this +# workload — counting `cudaStreamSynchronize` reports "no syncs" and is +# wrong. Count the D2H. +# 2. Kernel launches per decode step, before and after. Op count is the +# disease (9,131/token, median kernel 1.18 us), so this matters more than +# wall time. +# 3. Interleaved A-B-A-B ms/token with the monotonic drift stated. This box +# showed 3.3% drift across one run — larger than a real arm difference — +# so A-then-B is a fabricated comparison. +# 4. Bit-parity: the GPU-gated `kv_fp8_fused_is_bit_identical_to_cpu_exact` +# test, plus its negative control. +# +# RULES ENFORCED HERE, NOT ASSUMED +# * The box is shared with the ArcGraph chain. Every timing leg holds +# /root/.arc-bench.lock under flock; a contended number is a fabricated +# number. +# * Exclusivity is asserted BEFORE AND AFTER every leg from +# `nvidia-smi --query-compute-apps` (the lock file is not an occupancy +# signal — an 87 GB server has been resident with no lock at all, and a +# 77 GB server has run with the lock reading FREE). A neighbour appearing +# mid-leg aborts the run. +# * Environment failure exits 2. A failed measurement exits 1. They are not +# the same thing and must never be reported as the same thing. +# * Engagement is asserted before any null/neutral result is believed: if +# the arm that is supposed to use the fused kernel shows the same D2H +# count as the arm that is not, the run aborts rather than reporting "no +# difference". +set -uo pipefail + +REPO="${REPO:-/root/arc-wt/fp8}" +MODEL="${MODEL:-deepseek-ai/DeepSeek-V4-Flash}" +ARCH="${ARCH:-deepseekv4}" +ISQ="${ISQ:-qtip2}" +PROMPT_LEN="${PROMPT_LEN:-128}" +GEN_LEN="${GEN_LEN:-64}" +MAX_SEQ_LEN="${MAX_SEQ_LEN:-1024}" +OUT="${OUT:-/root/kvfp8-fused}" +LOCK="${LOCK:-/root/locks/bench.lock}" +BIN="${BIN:-${CARGO_TARGET_DIR:-$REPO/target}/release/mistralrs}" +REPS="${REPS:-2}" # A-B-A-B => REPS=2 + +mkdir -p "$OUT" + +envfail() { + echo "ENVFAIL: $*" >&2 + exit 2 +} +fail() { + echo "FAIL: $*" >&2 + exit 1 +} + +command -v nvidia-smi >/dev/null 2>&1 || envfail "no nvidia-smi" +command -v nsys >/dev/null 2>&1 || echo "WARN: no nsys; the D2H/launch counts will be skipped" >&2 +[ -x "$BIN" ] || envfail "no mistralrs binary at $BIN (build with --features 'cuda flash-attn')" + +# --------------------------------------------------------------------------- +# Exclusivity. Bracket every leg: a V4 load shows near-zero VRAM for most of a +# minute, so "compute-apps empty" sampled during a neighbour's load is +# indistinguishable from an idle box. +# --------------------------------------------------------------------------- +MYPID="" +assert_exclusive() { + local where="$1" + local apps + apps=$(nvidia-smi --query-compute-apps=pid --format=csv,noheader | tr -d ' ' | sort -u | tr '\n' ',') + apps="${apps%,}" + if [ -n "$MYPID" ]; then + [ "$apps" = "$MYPID" ] || envfail "[$where] compute-apps=[$apps], expected only $MYPID" + else + [ -z "$apps" ] || envfail "[$where] box not idle: compute-apps=[$apps]" + fi +} + +# --------------------------------------------------------------------------- +# One measured leg. $1 = arm label, $2 = ARC_KV_FP8_MODE value ("" => default, +# which is the fused kernel), $3 = log path, $4 = "nsys" to trace. +# --------------------------------------------------------------------------- +run_leg() { + local arm="$1" mode="$2" log="$3" trace="${4:-}" + local -a cmd=("$BIN" bench -m "$MODEL" -a "$ARCH" --isq "$ISQ" + --prompt-len "$PROMPT_LEN" --gen-len "$GEN_LEN" --max-seq-len "$MAX_SEQ_LEN") + + assert_exclusive "pre-$arm" + ( + if [ -n "$mode" ]; then export ARC_KV_FP8_MODE="$mode"; else unset ARC_KV_FP8_MODE; fi + unset ARC_GPU_ACT_QUANT + if [ "$trace" = "nsys" ]; then + nsys profile -t cuda -o "$OUT/$arm" --force-overwrite true \ + --cuda-memory-usage false "${cmd[@]}" + else + "${cmd[@]}" + fi + ) >"$log" 2>&1 & + local pid=$! + # `pgrep -f`/`pkill -f` would match this script's own command line; use the + # job pid we already hold instead. + MYPID="" + wait "$pid" + local rc=$? + MYPID="" + assert_exclusive "post-$arm" + [ $rc -eq 0 ] || fail "$arm leg exited $rc; see $log" +} + +# --------------------------------------------------------------------------- +# ms/token out of a bench log. Fails loudly rather than emitting an empty +# string that would silently become a "0.0 ms/token improvement". +# --------------------------------------------------------------------------- +decode_ms_per_token() { + local log="$1" + local tps + tps=$(grep -oE 'tok_per_s_decode[^0-9]*[0-9]+\.[0-9]+' "$log" | tail -1 | + grep -oE '[0-9]+\.[0-9]+$') + [ -n "$tps" ] || tps=$(grep -oiE 'decode[^0-9]*([0-9]+\.[0-9]+) *tok' "$log" | tail -1 | + grep -oE '[0-9]+\.[0-9]+') + [ -n "$tps" ] || fail "no decode throughput in $log — do not report a delta from a missing number" + awk -v t="$tps" 'BEGIN { printf "%.3f", 1000.0 / t }' +} + +# --------------------------------------------------------------------------- +# 0. Bit-parity + its negative control. Cheap, and it gates everything else: +# a faster kernel that stores different bytes is not a faster kernel. +# --------------------------------------------------------------------------- +echo "=== 0. bit-parity (GPU-gated test + D33 negative control) ===" +( + cd "$REPO" || envfail "no repo at $REPO" + cargo test -p mistralrs-core --features "cuda flash-attn" --lib \ + kv_fp8_fused_is_bit_identical_to_cpu_exact -- --nocapture --exact \ + models::dsv4_kv_fp8::tests::kv_fp8_fused_is_bit_identical_to_cpu_exact +) 2>&1 | tee "$OUT/parity.log" +prc=${PIPESTATUS[0]} +# `set -e` does not survive a pipe; the status is read explicitly. +[ "$prc" -eq 2 ] && envfail "parity test could not find a CUDA device" +[ "$prc" -eq 0 ] || fail "bit-parity FAILED — the fused kernel does not store what the CPU path stores" +grep -q "test result: ok. 1 passed" "$OUT/parity.log" || + fail "parity test reported no results — 'no failures' is not 'ran'" +echo "bit-parity: PASS" + +# --------------------------------------------------------------------------- +# 1 + 2. D2H count and launch count per decode step, both arms, under nsys. +# --------------------------------------------------------------------------- +if command -v nsys >/dev/null 2>&1; then + echo "=== 1+2. nsys: DtoH copies and kernel launches per step ===" + flock "$LOCK" bash -c "$(declare -f run_leg assert_exclusive envfail fail); \ + OUT='$OUT' BIN='$BIN' MODEL='$MODEL' ARCH='$ARCH' ISQ='$ISQ' \ + PROMPT_LEN='$PROMPT_LEN' GEN_LEN='$GEN_LEN' MAX_SEQ_LEN='$MAX_SEQ_LEN' \ + bash -c 'true'" || true + + for arm in before after; do + mode=""; [ "$arm" = "before" ] && mode="cpu" + ( + flock 9 || envfail "could not take the box lock $LOCK" + run_leg "$arm" "$mode" "$OUT/$arm.nsys.log" nsys + ) 9>"$LOCK" + nsys stats --report cuda_api_sum --format csv "$OUT/$arm.nsys-rep" \ + >"$OUT/$arm.api.csv" 2>"$OUT/$arm.api.err" || + echo "WARN: nsys stats cuda_api_sum failed for $arm" >&2 + nsys stats --report cuda_gpu_kern_sum --format csv "$OUT/$arm.nsys-rep" \ + >"$OUT/$arm.kern.csv" 2>>"$OUT/$arm.api.err" || + echo "WARN: nsys stats cuda_gpu_kern_sum failed for $arm" >&2 + + d2h=$(awk -F, '/cuMemcpyDtoHAsync_v2/ { gsub(/"/,"",$3); s+=$3 } END { print s+0 }' "$OUT/$arm.api.csv") + launches=$(awk -F, 'NR>1 { gsub(/"/,"",$3); s+=$3 } END { print s+0 }' "$OUT/$arm.kern.csv") + echo "$arm: cuMemcpyDtoHAsync_v2=$d2h kernel_launches=$launches (steps=$GEN_LEN)" + echo "$arm $d2h $launches" >>"$OUT/counts.txt" + done + + # Engagement (D18): the two arms MUST differ in D2H count. Equal counts mean + # the fused arm never engaged, and every later number is the same code twice. + b=$(awk '$1=="before"{print $2}' "$OUT/counts.txt") + a=$(awk '$1=="after"{print $2}' "$OUT/counts.txt") + [ -n "$b" ] && [ -n "$a" ] || fail "missing D2H counts — no results is not a null result" + [ "$a" -lt "$b" ] || fail "ENGAGEMENT: after-arm D2H ($a) is not below before-arm ($b); the fused kernel did not run" + awk -v b="$b" -v a="$a" -v n="$GEN_LEN" 'BEGIN { + printf "D2H per step: before %.2f -> after %.2f (removed %.2f/step)\n", b/n, a/n, (b-a)/n }' +fi + +# --------------------------------------------------------------------------- +# 3. Interleaved A-B-A-B ms/token. A = CPU round trip, B = fused kernel. +# The drift is reported alongside the delta, because on this box a 3.3% +# monotonic drift once exceeded the arm difference. +# --------------------------------------------------------------------------- +echo "=== 3. interleaved A-B-A-B ms/token ===" +: >"$OUT/ab.txt" +for i in $(seq 1 "$REPS"); do + for arm in A B; do + mode=""; [ "$arm" = "A" ] && mode="cpu" + log="$OUT/${arm}${i}.log" + ( + flock 9 || envfail "could not take the box lock $LOCK" + run_leg "${arm}${i}" "$mode" "$log" + ) 9>"$LOCK" + ms=$(decode_ms_per_token "$log") + echo "$arm $i $ms" >>"$OUT/ab.txt" + echo "${arm}${i} = $ms ms/token" + done +done + +awk ' + { v[$1""$2] = $3; arm[$1] += $3; n[$1]++ } + END { + a = arm["A"] / n["A"]; b = arm["B"] / n["B"]; + printf "A (CPU round trip) mean %.3f ms/token over %d\n", a, n["A"]; + printf "B (fused kernel) mean %.3f ms/token over %d\n", b, n["B"]; + printf "delta %.3f ms/token (%.2f%%)\n", a - b, 100.0 * (a - b) / a; + # Monotonic drift: same arm, first rep vs last rep. + if (n["A"] > 1) printf "drift within A %.2f%% (A1 %.3f -> A%d %.3f)\n", \ + 100.0 * (v["A" n["A"]] - v["A1"]) / v["A1"], v["A1"], n["A"], v["A" n["A"]]; + if (n["B"] > 1) printf "drift within B %.2f%% (B1 %.3f -> B%d %.3f)\n", \ + 100.0 * (v["B" n["B"]] - v["B1"]) / v["B1"], v["B1"], n["B"], v["B" n["B"]]; + print "REPORT THE DRIFT WITH THE DELTA. A delta smaller than the drift is not a result."; + }' "$OUT/ab.txt" + +echo "=== artifacts in $OUT ===" diff --git a/mistralrs-quant/kernels/arc_kvquant/arc_kvquant.cu b/mistralrs-quant/kernels/arc_kvquant/arc_kvquant.cu index 7afb89f6e..2f470e95e 100644 --- a/mistralrs-quant/kernels/arc_kvquant/arc_kvquant.cu +++ b/mistralrs-quant/kernels/arc_kvquant/arc_kvquant.cu @@ -36,11 +36,14 @@ // // 1. THIS CRATE COMPILES WITH `--use_fast_math` (mistralrs-quant/build.rs), so // a bare `a / b` here would be `div.approx.f32` while candle-kernels (no -// fast math) emits IEEE `div.rn.f32`, and `-ftz=true` would flush denormals -// candle keeps. Every float operation below is therefore an explicit -// `__f*_rn` intrinsic, which the CUDA Math API defines as IEEE-754 and -// unaffected by `-prec-div`/`-ftz`/`-fmad`. `fabsf`/`fmaxf` are avoided in -// favour of integer ops on the bit pattern for the same reason. +// fast math) emits IEEE `div.rn.f32`. The `__f*_rn` intrinsics are NOT a +// sufficient answer — checked against the emitted PTX, nvcc 13.1 rewrites +// them to `mul.rn.ftz.f32` / `div.rn.ftz.f32` / `add.rn.ftz.f32` under fast +// math, which flushes denormals candle keeps. Every float operation below is +// therefore INLINE PTX (`arc_fmul`/`arc_fadd`/`arc_fdiv`), which no +// optimisation flag rewrites; `fabsf`/`fmaxf` are avoided in favour of +// integer ops on the bit pattern for the same reason. The audit is one line: +// `nvcc -ptx ... | grep -c '\.ftz\.f32'` must be 0. // 2. The E4M3 rounding is a transcription of NVIDIA's // `__nv_cvt_double_to_fp8(x, __NV_SATFINITE, __NV_E4M3)` (cuda_fp8.hpp), // which is *also* what the Rust `float8` crate ports in `convert_to_fp8` @@ -72,6 +75,38 @@ namespace arc { // (inline `cvt.rn.bf16.f32`, immune to -ftz). // f32 -> f16 : likewise `__float2half` (inline `cvt.rn.f16.f32`). // --------------------------------------------------------------------------- +// --------------------------------------------------------------------------- +// IEEE f32 ops as inline PTX. +// +// MEASURED, not assumed: `--use_fast_math` (this crate's build.rs) rewrites the +// `__fmul_rn`/`__fadd_rn`/`__fdiv_rn` intrinsics to their `.ftz` forms — nvcc +// 13.1 emitted `mul.rn.ftz.f32` / `div.rn.ftz.f32` / `add.rn.ftz.f32` for them. +// The rounding mode survives, but denormal flush-to-zero does not match +// candle-kernels, which compiles WITHOUT fast math and emits `div.rn.f32`. +// +// The difference happens to be unobservable here (`scale` is floored at 1e-12, +// and E4M3's smallest denormal is 2^-9, so any f32 denormal is quantized to +// zero either way) — but bit-parity should not rest on a chain of reasoning +// about where denormals cannot reach. Inline PTX is never rewritten by +// optimisation flags, so these ops are IEEE non-ftz by construction, and +// `grep -c '\.ftz\.f32'` over the emitted PTX is a one-line audit of that. +// --------------------------------------------------------------------------- +__device__ __forceinline__ float arc_fmul(float a, float b) { + float r; + asm("mul.rn.f32 %0, %1, %2;" : "=f"(r) : "f"(a), "f"(b)); + return r; +} +__device__ __forceinline__ float arc_fadd(float a, float b) { + float r; + asm("add.rn.f32 %0, %1, %2;" : "=f"(r) : "f"(a), "f"(b)); + return r; +} +__device__ __forceinline__ float arc_fdiv(float a, float b) { + float r; + asm("div.rn.f32 %0, %1, %2;" : "=f"(r) : "f"(a), "f"(b)); + return r; +} + template __device__ __forceinline__ float arc_to_f32(T v); template <> __device__ __forceinline__ float arc_to_f32<__nv_bfloat16>(__nv_bfloat16 v) { return __bfloat162float(v); @@ -98,12 +133,12 @@ template <> __device__ __forceinline__ float arc_from_f32(float v) { retu // `affine(mul, add)` to `x * mul + add` with `mul`/`add` first narrowed to the // tensor dtype. `x * 1.0f` is exact, so the second affine is exactly an add, // and FMA contraction of `x * mul + 0.0f` is exactly the multiply. Both are -// therefore reproduced by one `__fmul_rn` and one `__fadd_rn`. +// therefore reproduced by one `arc_fmul` and one `arc_fadd` (inline PTX, see above). // --------------------------------------------------------------------------- __device__ __forceinline__ float arc_kv_block_scale(float amax) { const float inv_max = (float)(1.0 / ARC_KV_E4M3_MAX); const float eps = (float)1e-12; - return __fadd_rn(__fmul_rn(amax, inv_max), eps); + return arc_fadd(arc_fmul(amax, inv_max), eps); } // --------------------------------------------------------------------------- @@ -218,7 +253,7 @@ __global__ void arc_kv_fp8_quantize_kernel( for (int e = lane; e < block_w; e += ARC_KV_WARP) { const float v = arc_to_f32(krow[base + e]); - crow[base + e] = arc_f32_to_e4m3_code(__fdiv_rn(v, scale)); + crow[base + e] = arc_f32_to_e4m3_code(arc_fdiv(v, scale)); } // `amax` IS one of the block's own elements, so narrowing it back to the // activation dtype is exact - that is what lets dequant rebuild the @@ -270,7 +305,7 @@ __global__ void arc_kv_fp8_dequantize_kernel( const float amax = arc_to_f32(srow[rope_dim + blk]); const float scale = arc_kv_block_scale(amax); for (int e = lane; e < block_w; e += ARC_KV_WARP) { - orow[base + e] = arc_from_f32(__fmul_rn(lut[crow[base + e]], scale)); + orow[base + e] = arc_from_f32(arc_fmul(lut[crow[base + e]], scale)); } } @@ -419,7 +454,7 @@ __global__ void arc_kv_fp8_quantize_mutant_kernel( const float scale = arc_kv_block_scale(amax); for (int e = lane; e < block_w; e += ARC_KV_WARP) { const float v = arc_to_f32(krow[base + e]); - crow[base + e] = arc_f32_to_e4m3_code_truncating(__fdiv_rn(v, scale)); + crow[base + e] = arc_f32_to_e4m3_code_truncating(arc_fdiv(v, scale)); } if (lane == 0) { srow[rope_dim + blk] = arc_from_f32(amax); From 259369b2ee5dd9cad1dfb3c0249ac5596359f08c Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Tue, 18 Aug 2026 10:49:54 +0100 Subject: [PATCH 03/22] test(ArcKV/Fp8): exhaustive 2^32 E4M3 sweep + its negative control Runs on hardware with nvcc alone (no cargo build). Compares the kernel's transcription against NVIDIA's software reference (what the Rust float8 crate ports, hence what candle's CPU cast computes) and against the sm_90 hardware path, over ALL 2^32 f32 bit patterns. Measured on H200 / CUDA 13.1: visited 4,294,967,296 of 4,294,967,296 mismatch vs NVIDIA sw 0 mismatch sw vs hw 0 negative control 123,731,850 inputs caught (2.88%) The visited counter and the -DMUTANT=1 control exist because '0 mismatches' from a sweep that ran zero iterations is indistinguishable from a pass. Co-Authored-By: Claude Opus 5 (1M context) --- arc-tools/e4m3_exhaustive.cu | 137 +++++++++++++++++++++++++++++++++++ 1 file changed, 137 insertions(+) create mode 100644 arc-tools/e4m3_exhaustive.cu diff --git a/arc-tools/e4m3_exhaustive.cu b/arc-tools/e4m3_exhaustive.cu new file mode 100644 index 000000000..a93db1348 --- /dev/null +++ b/arc-tools/e4m3_exhaustive.cu @@ -0,0 +1,137 @@ +// EXHAUSTIVE proof that arc_kvquant.cu's E4M3 conversion IS NVIDIA's. +// +// Parent system: ArcQuant / TurboQuant (ArcInfer / ArcKV / Fp8). +// +// D33 — a check you have not seen fail is not a check. This sweeps ALL 2^32 +// f32 bit patterns (not a sample) and compares three implementations: +// +// arc::arc_f32_to_e4m3_code(x) the transcription the +// fused kernel ships +// __nv_cvt_double_to_fp8((double)x, SATFINITE, E4M3) NVIDIA's SOFTWARE +// reference — the routine +// the Rust `float8` crate +// ports as `convert_to_fp8`, +// hence exactly what +// `F8E4M3::from_f32`, and +// therefore candle's CPU +// cast, computes +// __nv_cvt_float_to_fp8(x, SATFINITE, E4M3) NVIDIA's HARDWARE path +// (cvt.rn.satfinite.e4m3x2.f32 +// on sm_89+) +// +// One mismatch against the software reference is a bug in the kernel. Software +// vs hardware is counted separately so the two can never be confused. +// +// Two guards make a green result mean something: +// * `visited` counts the values actually processed; a sweep that silently +// ran zero iterations would otherwise report "0 mismatches" and look +// identical to a pass. +// * `-DMUTANT=1` runs the same sweep against the deliberately-wrong +// truncating variant. If THAT also reports 0 mismatches the sweep is +// vacuous and the exit code says so. +// +// Build + run (needs only nvcc + a GPU, not the cargo build): +// nvcc -std=c++17 -O3 -arch=sm_90 --use_fast_math --expt-relaxed-constexpr \ +// -U__CUDA_NO_BFLOAT16_CONVERSIONS__ -DMUTANT=0 \ +// -I/mistralrs-quant/kernels/arc_kvquant e4m3_exhaustive.cu -o sweep0 +// Exit 0 pass, 1 mismatch, 2 environment/vacuous. +#ifndef MUTANT +#define MUTANT 0 +#endif + +#include +#include +#include +#include + +#include "arc_kvquant.cu" + +__global__ void sweep(unsigned long long base, unsigned long long n, + unsigned long long *mismatch_sw, + unsigned long long *mismatch_hw, unsigned *first_bad, + unsigned long long *visited) { + for (unsigned long long i = + blockIdx.x * (unsigned long long)blockDim.x + threadIdx.x; + i < n; i += (unsigned long long)gridDim.x * blockDim.x) { + const unsigned bits = (unsigned)(base + i); + const float x = __uint_as_float(bits); + const uint8_t mine = MUTANT ? arc::arc_f32_to_e4m3_code_truncating(x) + : arc::arc_f32_to_e4m3_code(x); + const uint8_t sw = + (uint8_t)__nv_cvt_double_to_fp8((double)x, __NV_SATFINITE, __NV_E4M3); + const uint8_t hw = + (uint8_t)__nv_cvt_float_to_fp8(x, __NV_SATFINITE, __NV_E4M3); + if (mine != sw) { + if (atomicAdd(mismatch_sw, 1ULL) == 0ULL) { + *first_bad = bits; + } + } + if (sw != hw) { + atomicAdd(mismatch_hw, 1ULL); + } + atomicAdd(visited, 1ULL); + } +} + +int main() { + unsigned long long *d_sw, *d_hw, *d_vis; + unsigned *d_first; + if (cudaMalloc(&d_sw, 8) != cudaSuccess || cudaMalloc(&d_hw, 8) != cudaSuccess || + cudaMalloc(&d_vis, 8) != cudaSuccess || + cudaMalloc(&d_first, 4) != cudaSuccess) { + printf("FATAL cudaMalloc\n"); + return 2; + } + cudaMemset(d_sw, 0, 8); + cudaMemset(d_hw, 0, 8); + cudaMemset(d_vis, 0, 8); + cudaMemset(d_first, 0, 4); + + const unsigned long long TOTAL = 1ULL << 32; + const unsigned long long CHUNK = 1ULL << 28; + for (unsigned long long base = 0; base < TOTAL; base += CHUNK) { + sweep<<<8192, 256>>>(base, CHUNK, d_sw, d_hw, d_first, d_vis); + cudaError_t e = cudaDeviceSynchronize(); + if (e != cudaSuccess) { + printf("FATAL cuda: %s\n", cudaGetErrorString(e)); + return 2; + } + } + + unsigned long long sw = 0, hw = 0, vis = 0; + unsigned first = 0; + cudaMemcpy(&sw, d_sw, 8, cudaMemcpyDeviceToHost); + cudaMemcpy(&hw, d_hw, 8, cudaMemcpyDeviceToHost); + cudaMemcpy(&vis, d_vis, 8, cudaMemcpyDeviceToHost); + cudaMemcpy(&first, d_first, 4, cudaMemcpyDeviceToHost); + + printf("MUTANT %d\n", MUTANT); + printf("visited %llu of %llu f32 bit patterns\n", vis, TOTAL); + if (vis != TOTAL) { + printf("FATAL the sweep did not visit every value; a 0 here means nothing\n"); + return 2; + } + printf("mismatch vs NVIDIA sw %llu\n", sw); + printf("mismatch sw vs hw %llu\n", hw); + + if (MUTANT) { + if (sw == 0) { + printf("FATAL VACUOUS: the deliberately-wrong kernel also matched on " + "every input, so this sweep cannot fail and proves nothing\n"); + return 2; + } + printf("NEGATIVE CONTROL OK: the wrong kernel is caught on %llu inputs " + "(%.2f%%)\n", + sw, 100.0 * (double)sw / (double)TOTAL); + return 0; + } + if (sw) { + float x; + memcpy(&x, &first, 4); + printf("first bad input 0x%08x (%.9g)\n", first, x); + return 1; + } + printf("RESULT: arc_f32_to_e4m3_code is bit-identical to NVIDIA's E4M3 " + "conversion on every one of the 2^32 f32 values.\n"); + return 0; +} From 112cbb4e1a37fbbe00fa1f2fa43b90c9eda18beb Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Tue, 18 Aug 2026 10:59:40 +0100 Subject: [PATCH 04/22] test(ArcLab): per-step D2H + launch counter, validated on a known answer The 43 blocking cuMemcpyDtoHAsync_v2 per V4 decode step are invisible to the obvious instrument (*Synchronize* is 0.0 calls/step), so this counts the copies themselves and pins the step count from the trace with two independent anchors rather than assuming it. Known-answer test on the recorded baseline trace (/root/budget-chain/nsys): D2H PER STEP 44.09 (BUDGET_V4_B1.md records 44) LAUNCHES PER STEP 9131.5 (records 9,131) 1,792 B x 8,144 = 43.09/step; 517,120 B x 189 = 1.00/step -- the exact two DtoH sizes the budget names. Instrument validated before being pointed at the new traces. Drops the earlier measure_kv_fp8_fused.sh: it was written before the box paths were known and its exclusivity check had a defect (it cleared the pid it was meant to compare against). Replaced by the scripts actually run on the box. Co-Authored-By: Claude Opus 5 (1M context) --- arc-tools/kv_fp8_count_per_step.py | 102 +++++++++++++ arc-tools/measure_kv_fp8_fused.sh | 223 ----------------------------- 2 files changed, 102 insertions(+), 223 deletions(-) create mode 100644 arc-tools/kv_fp8_count_per_step.py delete mode 100755 arc-tools/measure_kv_fp8_fused.sh diff --git a/arc-tools/kv_fp8_count_per_step.py b/arc-tools/kv_fp8_count_per_step.py new file mode 100644 index 000000000..224945b22 --- /dev/null +++ b/arc-tools/kv_fp8_count_per_step.py @@ -0,0 +1,102 @@ +#!/usr/bin/env python3 +"""Per-decode-step D2H and kernel-launch counts from an nsys sqlite. + +Parent system: ArcLab (measuring ArcInfer / ArcKV / Fp8). + +WHY THIS EXISTS + The 43 blocking `cuMemcpyDtoHAsync_v2` per V4 decode step are invisible to + the obvious instrument: `*Synchronize*` is 0.0 calls/step in this workload, + so anyone counting `cudaStreamSynchronize` concludes "no syncs" and is + wrong. This counts the copies themselves. + +THE STEP COUNT IS PINNED FROM THE TRACE, NOT ASSUMED + A per-step number divided by a guessed step count is a fabricated number. + Two independent anchors must agree: + A. the once-per-step logits readback -- the largest DtoH size, whose + occurrence count IS the number of steps in the window; + B. structural agreement -- kernels whose counts are integer multiples of + that same step count (the 43-layer loop makes many). Fewer than three + such kernels means the window is not steady-state decode and the + script exits 2. + +Exits 2 on environment/analysis failure, never 1. + +Usage: kv_fp8_count_per_step.py [tail_seconds] +""" +import collections +import sqlite3 +import sys + + +def die(msg): + print(f"FATAL {msg}", file=sys.stderr) + sys.exit(2) + + +def main(): + if len(sys.argv) < 2: + die("usage: kv_fp8_count_per_step.py [tail_seconds]") + tail_s = float(sys.argv[2]) if len(sys.argv) > 2 else 20.0 + + db = sqlite3.connect(sys.argv[1]) + cur = db.cursor() + strings = {i: v for i, v in cur.execute("SELECT id, value FROM StringIds")} + KT = "CUPTI_ACTIVITY_KIND_KERNEL" + MT = "CUPTI_ACTIVITY_KIND_MEMCPY" + + end = cur.execute(f"SELECT MAX(end) FROM {KT}").fetchone()[0] + if end is None: + die("no kernels in the trace") + win0 = end - int(tail_s * 1e9) + + # copyKind 2 == device-to-host. + sizes = list(cur.execute( + f"SELECT bytes, COUNT(*) FROM {MT} WHERE start>={win0} AND copyKind=2 " + f"GROUP BY bytes ORDER BY COUNT(*) DESC")) + if not sizes: + die("no DtoH copies in the window") + total_d2h = sum(c for _, c in sizes) + + # ---- Anchor A: the logits readback, once per decode step. + big = [(b, c) for b, c in sizes if b > 100_000] + if not big: + die("no large (>100 kB) DtoH found; cannot pin the step count") + steps = max(c for _, c in big) + if steps < 20: + die(f"only {steps} steps in the window; widen tail_seconds") + + kc = collections.Counter() + for (n,) in cur.execute(f"SELECT shortName FROM {KT} WHERE start>={win0}"): + kc[strings.get(n, str(n))] += 1 + if not kc: + die("no kernels in the window") + + # ---- Anchor B: kernels whose counts are integer multiples of `steps`. + agreeing = [] + for name, c in kc.items(): + k = c / steps + if k >= 0.98 and abs(k - round(k)) <= 0.02 * max(round(k), 1): + agreeing.append((name, c, round(k))) + if len(agreeing) < 3: + die(f"step count {steps} unconfirmed: only {len(agreeing)} kernels are " + f"integer multiples of it; the window is probably not steady-state " + f"decode") + + total_kern = sum(kc.values()) + print(f"window last {tail_s:.0f} s") + print(f"steps (logits anchor) {steps}") + print(f"steps confirmed by {len(agreeing)} kernels at integer multiples") + print(f"D2H total {total_d2h}") + print(f"D2H PER STEP {total_d2h / steps:.2f}") + print(f"kernels total {total_kern}") + print(f"LAUNCHES PER STEP {total_kern / steps:.1f}") + print("DtoH by size:") + for b, c in sizes[:8]: + print(f" {b:>12,d} B x {c:>9,d} {c / steps:8.2f}/step") + print("anchor-B kernels (count = k x steps):") + for name, c, k in sorted(agreeing, key=lambda t: -t[1])[:6]: + print(f" {c:>9,d} = {k:>4d} x steps {name[:52]}") + + +if __name__ == "__main__": + main() diff --git a/arc-tools/measure_kv_fp8_fused.sh b/arc-tools/measure_kv_fp8_fused.sh deleted file mode 100755 index fcf03815a..000000000 --- a/arc-tools/measure_kv_fp8_fused.sh +++ /dev/null @@ -1,223 +0,0 @@ -#!/usr/bin/env bash -# ArcKV/Fp8 — measure the fused E4M3 quantize+dequantize kernel against the -# CPU round trip it replaces. Runs ON THE BOX. Nothing here is a proxy: every -# number is produced by the real binary on the real model. -# -# WHAT IT PRODUCES (the four numbers the change is judged on) -# 1. `cuMemcpyDtoHAsync_v2` calls per decode step, before and after. -# The 44/step are the thing being removed and the reason CUDA graph -# capture is impossible. NOTE `*Synchronize*` is 0.0 calls/step in this -# workload — counting `cudaStreamSynchronize` reports "no syncs" and is -# wrong. Count the D2H. -# 2. Kernel launches per decode step, before and after. Op count is the -# disease (9,131/token, median kernel 1.18 us), so this matters more than -# wall time. -# 3. Interleaved A-B-A-B ms/token with the monotonic drift stated. This box -# showed 3.3% drift across one run — larger than a real arm difference — -# so A-then-B is a fabricated comparison. -# 4. Bit-parity: the GPU-gated `kv_fp8_fused_is_bit_identical_to_cpu_exact` -# test, plus its negative control. -# -# RULES ENFORCED HERE, NOT ASSUMED -# * The box is shared with the ArcGraph chain. Every timing leg holds -# /root/.arc-bench.lock under flock; a contended number is a fabricated -# number. -# * Exclusivity is asserted BEFORE AND AFTER every leg from -# `nvidia-smi --query-compute-apps` (the lock file is not an occupancy -# signal — an 87 GB server has been resident with no lock at all, and a -# 77 GB server has run with the lock reading FREE). A neighbour appearing -# mid-leg aborts the run. -# * Environment failure exits 2. A failed measurement exits 1. They are not -# the same thing and must never be reported as the same thing. -# * Engagement is asserted before any null/neutral result is believed: if -# the arm that is supposed to use the fused kernel shows the same D2H -# count as the arm that is not, the run aborts rather than reporting "no -# difference". -set -uo pipefail - -REPO="${REPO:-/root/arc-wt/fp8}" -MODEL="${MODEL:-deepseek-ai/DeepSeek-V4-Flash}" -ARCH="${ARCH:-deepseekv4}" -ISQ="${ISQ:-qtip2}" -PROMPT_LEN="${PROMPT_LEN:-128}" -GEN_LEN="${GEN_LEN:-64}" -MAX_SEQ_LEN="${MAX_SEQ_LEN:-1024}" -OUT="${OUT:-/root/kvfp8-fused}" -LOCK="${LOCK:-/root/locks/bench.lock}" -BIN="${BIN:-${CARGO_TARGET_DIR:-$REPO/target}/release/mistralrs}" -REPS="${REPS:-2}" # A-B-A-B => REPS=2 - -mkdir -p "$OUT" - -envfail() { - echo "ENVFAIL: $*" >&2 - exit 2 -} -fail() { - echo "FAIL: $*" >&2 - exit 1 -} - -command -v nvidia-smi >/dev/null 2>&1 || envfail "no nvidia-smi" -command -v nsys >/dev/null 2>&1 || echo "WARN: no nsys; the D2H/launch counts will be skipped" >&2 -[ -x "$BIN" ] || envfail "no mistralrs binary at $BIN (build with --features 'cuda flash-attn')" - -# --------------------------------------------------------------------------- -# Exclusivity. Bracket every leg: a V4 load shows near-zero VRAM for most of a -# minute, so "compute-apps empty" sampled during a neighbour's load is -# indistinguishable from an idle box. -# --------------------------------------------------------------------------- -MYPID="" -assert_exclusive() { - local where="$1" - local apps - apps=$(nvidia-smi --query-compute-apps=pid --format=csv,noheader | tr -d ' ' | sort -u | tr '\n' ',') - apps="${apps%,}" - if [ -n "$MYPID" ]; then - [ "$apps" = "$MYPID" ] || envfail "[$where] compute-apps=[$apps], expected only $MYPID" - else - [ -z "$apps" ] || envfail "[$where] box not idle: compute-apps=[$apps]" - fi -} - -# --------------------------------------------------------------------------- -# One measured leg. $1 = arm label, $2 = ARC_KV_FP8_MODE value ("" => default, -# which is the fused kernel), $3 = log path, $4 = "nsys" to trace. -# --------------------------------------------------------------------------- -run_leg() { - local arm="$1" mode="$2" log="$3" trace="${4:-}" - local -a cmd=("$BIN" bench -m "$MODEL" -a "$ARCH" --isq "$ISQ" - --prompt-len "$PROMPT_LEN" --gen-len "$GEN_LEN" --max-seq-len "$MAX_SEQ_LEN") - - assert_exclusive "pre-$arm" - ( - if [ -n "$mode" ]; then export ARC_KV_FP8_MODE="$mode"; else unset ARC_KV_FP8_MODE; fi - unset ARC_GPU_ACT_QUANT - if [ "$trace" = "nsys" ]; then - nsys profile -t cuda -o "$OUT/$arm" --force-overwrite true \ - --cuda-memory-usage false "${cmd[@]}" - else - "${cmd[@]}" - fi - ) >"$log" 2>&1 & - local pid=$! - # `pgrep -f`/`pkill -f` would match this script's own command line; use the - # job pid we already hold instead. - MYPID="" - wait "$pid" - local rc=$? - MYPID="" - assert_exclusive "post-$arm" - [ $rc -eq 0 ] || fail "$arm leg exited $rc; see $log" -} - -# --------------------------------------------------------------------------- -# ms/token out of a bench log. Fails loudly rather than emitting an empty -# string that would silently become a "0.0 ms/token improvement". -# --------------------------------------------------------------------------- -decode_ms_per_token() { - local log="$1" - local tps - tps=$(grep -oE 'tok_per_s_decode[^0-9]*[0-9]+\.[0-9]+' "$log" | tail -1 | - grep -oE '[0-9]+\.[0-9]+$') - [ -n "$tps" ] || tps=$(grep -oiE 'decode[^0-9]*([0-9]+\.[0-9]+) *tok' "$log" | tail -1 | - grep -oE '[0-9]+\.[0-9]+') - [ -n "$tps" ] || fail "no decode throughput in $log — do not report a delta from a missing number" - awk -v t="$tps" 'BEGIN { printf "%.3f", 1000.0 / t }' -} - -# --------------------------------------------------------------------------- -# 0. Bit-parity + its negative control. Cheap, and it gates everything else: -# a faster kernel that stores different bytes is not a faster kernel. -# --------------------------------------------------------------------------- -echo "=== 0. bit-parity (GPU-gated test + D33 negative control) ===" -( - cd "$REPO" || envfail "no repo at $REPO" - cargo test -p mistralrs-core --features "cuda flash-attn" --lib \ - kv_fp8_fused_is_bit_identical_to_cpu_exact -- --nocapture --exact \ - models::dsv4_kv_fp8::tests::kv_fp8_fused_is_bit_identical_to_cpu_exact -) 2>&1 | tee "$OUT/parity.log" -prc=${PIPESTATUS[0]} -# `set -e` does not survive a pipe; the status is read explicitly. -[ "$prc" -eq 2 ] && envfail "parity test could not find a CUDA device" -[ "$prc" -eq 0 ] || fail "bit-parity FAILED — the fused kernel does not store what the CPU path stores" -grep -q "test result: ok. 1 passed" "$OUT/parity.log" || - fail "parity test reported no results — 'no failures' is not 'ran'" -echo "bit-parity: PASS" - -# --------------------------------------------------------------------------- -# 1 + 2. D2H count and launch count per decode step, both arms, under nsys. -# --------------------------------------------------------------------------- -if command -v nsys >/dev/null 2>&1; then - echo "=== 1+2. nsys: DtoH copies and kernel launches per step ===" - flock "$LOCK" bash -c "$(declare -f run_leg assert_exclusive envfail fail); \ - OUT='$OUT' BIN='$BIN' MODEL='$MODEL' ARCH='$ARCH' ISQ='$ISQ' \ - PROMPT_LEN='$PROMPT_LEN' GEN_LEN='$GEN_LEN' MAX_SEQ_LEN='$MAX_SEQ_LEN' \ - bash -c 'true'" || true - - for arm in before after; do - mode=""; [ "$arm" = "before" ] && mode="cpu" - ( - flock 9 || envfail "could not take the box lock $LOCK" - run_leg "$arm" "$mode" "$OUT/$arm.nsys.log" nsys - ) 9>"$LOCK" - nsys stats --report cuda_api_sum --format csv "$OUT/$arm.nsys-rep" \ - >"$OUT/$arm.api.csv" 2>"$OUT/$arm.api.err" || - echo "WARN: nsys stats cuda_api_sum failed for $arm" >&2 - nsys stats --report cuda_gpu_kern_sum --format csv "$OUT/$arm.nsys-rep" \ - >"$OUT/$arm.kern.csv" 2>>"$OUT/$arm.api.err" || - echo "WARN: nsys stats cuda_gpu_kern_sum failed for $arm" >&2 - - d2h=$(awk -F, '/cuMemcpyDtoHAsync_v2/ { gsub(/"/,"",$3); s+=$3 } END { print s+0 }' "$OUT/$arm.api.csv") - launches=$(awk -F, 'NR>1 { gsub(/"/,"",$3); s+=$3 } END { print s+0 }' "$OUT/$arm.kern.csv") - echo "$arm: cuMemcpyDtoHAsync_v2=$d2h kernel_launches=$launches (steps=$GEN_LEN)" - echo "$arm $d2h $launches" >>"$OUT/counts.txt" - done - - # Engagement (D18): the two arms MUST differ in D2H count. Equal counts mean - # the fused arm never engaged, and every later number is the same code twice. - b=$(awk '$1=="before"{print $2}' "$OUT/counts.txt") - a=$(awk '$1=="after"{print $2}' "$OUT/counts.txt") - [ -n "$b" ] && [ -n "$a" ] || fail "missing D2H counts — no results is not a null result" - [ "$a" -lt "$b" ] || fail "ENGAGEMENT: after-arm D2H ($a) is not below before-arm ($b); the fused kernel did not run" - awk -v b="$b" -v a="$a" -v n="$GEN_LEN" 'BEGIN { - printf "D2H per step: before %.2f -> after %.2f (removed %.2f/step)\n", b/n, a/n, (b-a)/n }' -fi - -# --------------------------------------------------------------------------- -# 3. Interleaved A-B-A-B ms/token. A = CPU round trip, B = fused kernel. -# The drift is reported alongside the delta, because on this box a 3.3% -# monotonic drift once exceeded the arm difference. -# --------------------------------------------------------------------------- -echo "=== 3. interleaved A-B-A-B ms/token ===" -: >"$OUT/ab.txt" -for i in $(seq 1 "$REPS"); do - for arm in A B; do - mode=""; [ "$arm" = "A" ] && mode="cpu" - log="$OUT/${arm}${i}.log" - ( - flock 9 || envfail "could not take the box lock $LOCK" - run_leg "${arm}${i}" "$mode" "$log" - ) 9>"$LOCK" - ms=$(decode_ms_per_token "$log") - echo "$arm $i $ms" >>"$OUT/ab.txt" - echo "${arm}${i} = $ms ms/token" - done -done - -awk ' - { v[$1""$2] = $3; arm[$1] += $3; n[$1]++ } - END { - a = arm["A"] / n["A"]; b = arm["B"] / n["B"]; - printf "A (CPU round trip) mean %.3f ms/token over %d\n", a, n["A"]; - printf "B (fused kernel) mean %.3f ms/token over %d\n", b, n["B"]; - printf "delta %.3f ms/token (%.2f%%)\n", a - b, 100.0 * (a - b) / a; - # Monotonic drift: same arm, first rep vs last rep. - if (n["A"] > 1) printf "drift within A %.2f%% (A1 %.3f -> A%d %.3f)\n", \ - 100.0 * (v["A" n["A"]] - v["A1"]) / v["A1"], v["A1"], n["A"], v["A" n["A"]]; - if (n["B"] > 1) printf "drift within B %.2f%% (B1 %.3f -> B%d %.3f)\n", \ - 100.0 * (v["B" n["B"]] - v["B1"]) / v["B1"], v["B1"], n["B"], v["B" n["B"]]; - print "REPORT THE DRIFT WITH THE DELTA. A delta smaller than the drift is not a result."; - }' "$OUT/ab.txt" - -echo "=== artifacts in $OUT ===" From 41a9343dfd891b0e0c6580a468f87e107d1a8006 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Tue, 18 Aug 2026 11:03:25 +0100 Subject: [PATCH 05/22] test(ArcKV/Fp8): nsys count harness with lock + VRAM gate MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Counts survive contention (nsys traces only this process), but V4 is ~79 GB of the H200's 143 GB, so holding the bench lock is not enough — the previous holder's process can still be resident. Gates on nvidia-smi free memory and exits 2 on OOM or a missing report rather than reporting a partial trace. Co-Authored-By: Claude Opus 5 (1M context) --- arc-tools/kv_fp8_nsys_ab.sh | 67 +++++++++++++++++++++++++++++++++++++ 1 file changed, 67 insertions(+) create mode 100755 arc-tools/kv_fp8_nsys_ab.sh diff --git a/arc-tools/kv_fp8_nsys_ab.sh b/arc-tools/kv_fp8_nsys_ab.sh new file mode 100755 index 000000000..6d4d40fbe --- /dev/null +++ b/arc-tools/kv_fp8_nsys_ab.sh @@ -0,0 +1,67 @@ +#!/bin/bash +# ArcKV/Fp8 — nsys legs: `cuMemcpyDtoHAsync_v2` copies and kernel launches PER +# DECODE STEP, before (CPU round trip) and after (fused device kernel). +# +# These are COUNTS. Box contention changes wall-clock, not how many D2H calls a +# decode step makes, so a contended run still yields a valid count — but the +# card must physically have room: V4 is ~79 GB of the H200's 143 GB, so two of +# these do not fit, and an OOM mid-trace is an environment failure (exit 2), not +# a result. +# +# Counted with arc-tools/kv_fp8_count_per_step.py, which pins the step count +# from the trace and was first validated against the recorded baseline (it +# reproduces 44.09 D2H/step and 9,131.5 launches/step). +set -u +OUT=${OUT:-/root/kvfp8} +BASE=${BASE:-/root/budget-chain} +BIN=${BIN:-/root/arc-wt/fp8-target/release/mistralrs} +DUR=${DUR:-280} +GEN=${GEN:-3000} +export PATH=/usr/local/cuda-13.1/bin:$PATH +export LD_LIBRARY_PATH=/usr/local/cuda/compat:${LD_LIBRARY_PATH:-} +unset ARC_TIME_DECODE V4_STATS V4_NAN_DEBUG V4_TRACE ARC_PROFILE ARC_GPU_ACT_QUANT +[ -x "$BIN" ] || { echo "FATAL_NO_BINARY $BIN"; exit 2; } +command -v nsys >/dev/null || { echo "FATAL_NO_NSYS"; exit 2; } + +# Holding the bench lock is not the same as the previous holder's process having +# exited; gate on what the card actually reports. +wait_for_vram() { + local need=${1:-100000} i=0 free + while :; do + free=$(nvidia-smi --query-gpu=memory.free --format=csv,noheader,nounits) + [ "${free:-0}" -ge "$need" ] && { echo "vram ok: ${free} MiB free"; return 0; } + i=$((i + 1)) + [ "$i" -gt 240 ] && { echo "FATAL_NO_VRAM free=${free}MiB after 20 min"; exit 2; } + sleep 5 + done +} + +for arm in before:cpu after:fused; do + A=${arm%%:*} + M=${arm##*:} + echo "=== leg $A (ARC_KV_FP8_MODE=$M) $(date -u +%T) ===" + wait_for_vram 100000 + rm -f "$OUT/$A".nsys-rep "$OUT/$A".sqlite + ARC_KV_FP8_MODE=$M nsys profile --trace=cuda --sample=none --cpuctxsw=none \ + --cuda-memory-usage=false --duration="$DUR" --kill=sigterm \ + --force-overwrite=true --output="$OUT/$A" \ + "$BIN" bench -m "$BASE/src" -a deepseekv4 \ + --from-uqff "$BASE/uqff/qtip2-0.uqff" --max-seqs 1 --prefix-cache-n 0 \ + --prompt-len 64 --gen-len "$GEN" --iterations 1 --warmup 0 \ + >"$OUT/$A.nsys.log" 2>&1 + RC=$? + echo "NSYS_RC=$RC" + # 143 = SIGTERM, which is how --kill ends the window and is EXPECTED. + [ "$RC" = "139" ] && { echo "FATAL_TARGET_SIGSEGV"; exit 2; } + grep -qi "out of memory" "$OUT/$A.nsys.log" && { echo "FATAL_OOM $A"; exit 2; } + [ -f "$OUT/$A.nsys-rep" ] || { echo "FATAL_NO_REPORT $A"; exit 2; } + nsys stats --force-export=true --report cuda_api_sum --format csv \ + --output "$OUT/${A}_api" "$OUT/$A.nsys-rep" >/dev/null 2>&1 + [ -f "$OUT/$A.sqlite" ] || { echo "FATAL_NO_SQLITE $A"; exit 2; } +done + +echo "=== ANALYSIS ===" +for A in before after; do + echo "--- $A ---" + python3 "$OUT/count.py" "$OUT/$A.sqlite" 20 || exit 2 +done From ad75bdc29ace43607d707a744a4474845b0c1fc7 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Tue, 18 Aug 2026 11:06:20 +0100 Subject: [PATCH 06/22] test(ArcLab): count CUDA API calls per step too, incl. the invisible sync MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Adds cuMemcpyDtoHAsync_v2 / cuLaunchKernel / cuMemAllocAsync / cuMemFreeAsync per step, and prints ALL *Synchronize* alongside — it reads 0.00/step, which is exactly why counting syncs on this workload finds nothing and concludes wrongly. Known-answer re-test on the recorded baseline trace now reproduces every headline in BUDGET_V4_B1.md: 44.09/step cuMemcpyDtoHAsync_v2 (recorded 44) 11436.33 + 11436.23 alloc/free/step (recorded 11,436 each) 2818.31/step cuMemcpyHtoDAsync_v2 (recorded 2,818) 0.00/step ALL *Synchronize* (recorded 0.0) 9131.5 launches/step (recorded 9,131) Co-Authored-By: Claude Opus 5 (1M context) --- arc-tools/kv_fp8_count_per_step.py | 23 +++++++++++++++++++++++ 1 file changed, 23 insertions(+) diff --git a/arc-tools/kv_fp8_count_per_step.py b/arc-tools/kv_fp8_count_per_step.py index 224945b22..39bb43314 100644 --- a/arc-tools/kv_fp8_count_per_step.py +++ b/arc-tools/kv_fp8_count_per_step.py @@ -97,6 +97,29 @@ def main(): for name, c, k in sorted(agreeing, key=lambda t: -t[1])[:6]: print(f" {c:>9,d} = {k:>4d} x steps {name[:52]}") + # ---- CUDA API calls per step. `cuMemcpyDtoHAsync_v2` is the one that + # matters here and it is NOT visible as a "sync": `*Synchronize*` is 0.0 + # calls/step in this workload. An "async" copy costing ~109 us of host time + # is a pageable-memory staged copy, i.e. blocking, and a CUDA graph cannot + # record one. + api = collections.Counter() + try: + for (n,) in cur.execute( + "SELECT nameId FROM CUPTI_ACTIVITY_KIND_RUNTIME WHERE start>=?", + (win0,)): + api[strings.get(n, str(n))] += 1 + except sqlite3.Error: + print("(no CUPTI_ACTIVITY_KIND_RUNTIME table; API counts unavailable)") + return + watch = ("cuMemcpyDtoHAsync_v2", "cuLaunchKernel", "cuMemAllocAsync", + "cuMemFreeAsync", "cuMemcpyHtoDAsync_v2", "cuMemsetD8Async") + print("CUDA API per step:") + for name in watch: + print(f" {api.get(name, 0) / steps:10.2f}/step {name}") + syncs = sum(c for n, c in api.items() if "Synchronize" in n) + print(f" {syncs / steps:10.2f}/step ALL *Synchronize* " + f"(this is why counting syncs finds nothing)") + if __name__ == "__main__": main() From 02c38946b4cad0c750ddaa7d0f9fe24cf67803cb Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Tue, 18 Aug 2026 11:38:48 +0100 Subject: [PATCH 07/22] fix(ArcLab): hold the bench lock across GPU work ONLY, not report export MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Measured on the box: a chain wrapping its whole pipeline in one flock held /root/locks/bench.lock with 13-17 waiters queued while the GPU read 0 %, 0 MiB, 78.28 W. nsys report export and the per-step counting are pure CPU on files already written, and must not hold the card. The script now takes the lock itself, per leg, around the VRAM wait and the traced run, and releases it the instant the bench exits — and says so in the output so the release time is auditable. Do not wrap it in flock. Co-Authored-By: Claude Opus 5 (1M context) --- arc-tools/kv_fp8_nsys_ab.sh | 18 ++++++++++++++++-- 1 file changed, 16 insertions(+), 2 deletions(-) diff --git a/arc-tools/kv_fp8_nsys_ab.sh b/arc-tools/kv_fp8_nsys_ab.sh index 6d4d40fbe..362a2b946 100755 --- a/arc-tools/kv_fp8_nsys_ab.sh +++ b/arc-tools/kv_fp8_nsys_ab.sh @@ -11,11 +11,20 @@ # Counted with arc-tools/kv_fp8_count_per_step.py, which pins the step count # from the trace and was first validated against the recorded baseline (it # reproduces 44.09 D2H/step and 9,131.5 launches/step). +# +# LOCK DISCIPLINE — this script takes /root/locks/bench.lock ITSELF, per leg, +# and holds it ONLY across the traced run. Do NOT wrap this script in `flock`: +# that holds the card through nsys report export and the counting, both of which +# are pure CPU on files already written. Measured on this box, a chain doing +# exactly that held the lock with 13-17 waiters queued while the GPU read +# `0 %, 0 MiB, 78.28 W`. set -u OUT=${OUT:-/root/kvfp8} BASE=${BASE:-/root/budget-chain} BIN=${BIN:-/root/arc-wt/fp8-target/release/mistralrs} DUR=${DUR:-280} +LOCK=${LOCK:-/root/locks/bench.lock} +LOCKWAIT=${LOCKWAIT:-5400} GEN=${GEN:-3000} export PATH=/usr/local/cuda-13.1/bin:$PATH export LD_LIBRARY_PATH=/usr/local/cuda/compat:${LD_LIBRARY_PATH:-} @@ -36,12 +45,16 @@ wait_for_vram() { done } +exec 9>"$LOCK" for arm in before:cpu after:fused; do A=${arm%%:*} M=${arm##*:} echo "=== leg $A (ARC_KV_FP8_MODE=$M) $(date -u +%T) ===" - wait_for_vram 100000 rm -f "$OUT/$A".nsys-rep "$OUT/$A".sqlite + # LOCK HELD ONLY HERE: VRAM wait + the traced run. Released the instant the + # bench exits, before any report export. + flock -w "$LOCKWAIT" 9 || { echo "FATAL_LOCK_TIMEOUT after ${LOCKWAIT}s"; exit 2; } + wait_for_vram 100000 ARC_KV_FP8_MODE=$M nsys profile --trace=cuda --sample=none --cpuctxsw=none \ --cuda-memory-usage=false --duration="$DUR" --kill=sigterm \ --force-overwrite=true --output="$OUT/$A" \ @@ -50,7 +63,8 @@ for arm in before:cpu after:fused; do --prompt-len 64 --gen-len "$GEN" --iterations 1 --warmup 0 \ >"$OUT/$A.nsys.log" 2>&1 RC=$? - echo "NSYS_RC=$RC" + flock -u 9 + echo "NSYS_RC=$RC (lock released $(date -u +%T); export below is CPU-only)" # 143 = SIGTERM, which is how --kill ends the window and is EXPECTED. [ "$RC" = "139" ] && { echo "FATAL_TARGET_SIGSEGV"; exit 2; } grep -qi "out of memory" "$OUT/$A.nsys.log" && { echo "FATAL_OOM $A"; exit 2; } From e57b95933ac60441ab0eb5e89a4f61791ce1d56c Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Tue, 18 Aug 2026 11:50:00 +0100 Subject: [PATCH 08/22] test(ArcKV/Fp8): interleaved A/B driver with per-leg lock + Xid check MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Lock covers exactly one leg (model load, which allocates, plus the timed run) and is released before parsing. Interleaves A-B-A-B and prints the within-arm drift beside the delta, because this box once drifted 3.3% monotonically — more than the arm difference. On a bench failure it dumps 'dmesg | grep -i xid' first: this box carries ~1,485 Xid GPU faults (ECC uncorrected 0, so not memory corruption), and a process killed by one dies with no error line. Distinguishing the box's fault from the code's is the difference between fixing a bug and chasing one that does not exist. Co-Authored-By: Claude Opus 5 (1M context) --- arc-tools/kv_fp8_ab.sh | 83 ++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 83 insertions(+) create mode 100755 arc-tools/kv_fp8_ab.sh diff --git a/arc-tools/kv_fp8_ab.sh b/arc-tools/kv_fp8_ab.sh new file mode 100755 index 000000000..876d5042c --- /dev/null +++ b/arc-tools/kv_fp8_ab.sh @@ -0,0 +1,83 @@ +#!/usr/bin/env bash +# ArcKV/Fp8 — interleaved A-B-A-B decode timing for the fused E4M3 kernel. +# A = ARC_KV_FP8_MODE=cpu (43 blocking D2H per step) +# B = ARC_KV_FP8_MODE=fused (one device kernel, 0 D2H) +# +# A-then-B is banned on this box: a 3.3% monotonic drift once exceeded the arm +# difference. The drift within each arm is printed next to the delta, and a +# delta smaller than the drift is not a result. +# +# LOCK DISCIPLINE. The lock covers exactly one thing: a leg (model load, which +# allocates GPU memory, plus the timed run). It is taken immediately before and +# released immediately after — NOT held across parsing, and NOT wrapped around +# the whole script. Six chains share this H200; one holding the lock through +# post-processing starved the trellis chain for 40 minutes while the card read +# 0 %, 0 MiB, 78 W. +set -u +OUT=${OUT:-/root/kvfp8} +BASE=${BASE:-/root/budget-chain} +BIN=${BIN:-/root/arc-wt/fp8-target/release/mistralrs} +LOCK=${LOCK:-/root/locks/bench.lock} +LOCKWAIT=${LOCKWAIT:-3600} +REPS=${REPS:-2} +GEN=${GEN:-260} +export PATH=/usr/local/cuda-13.1/bin:$PATH +export LD_LIBRARY_PATH=/usr/local/cuda/compat:${LD_LIBRARY_PATH:-} +unset ARC_TIME_DECODE V4_STATS V4_NAN_DEBUG V4_TRACE ARC_PROFILE ARC_GPU_ACT_QUANT +[ -x "$BIN" ] || { echo "FATAL_NO_BINARY $BIN"; exit 2; } + +exec 9>"$LOCK" + +leg() { # $1 = tag, $2 = mode + local tag=$1 mode=$2 log="$OUT/$1.log" free rc ms + flock -w "$LOCKWAIT" 9 || { echo "FATAL_LOCK_TIMEOUT ${LOCKWAIT}s"; exit 2; } + # Holding the lock is not the same as the card being free: V4 is ~79 GB of + # 143 GB and the previous holder's process can outlive its script. + for _ in $(seq 1 240); do + free=$(nvidia-smi --query-gpu=memory.free --format=csv,noheader,nounits) + [ "${free:-0}" -ge 100000 ] && break + sleep 5 + done + [ "${free:-0}" -ge 100000 ] || { echo "FATAL_NO_VRAM ${free}MiB"; exit 2; } + + ARC_KV_FP8_MODE=$mode "$BIN" bench -m "$BASE/src" -a deepseekv4 \ + --from-uqff "$BASE/uqff/qtip2-0.uqff" --max-seqs 1 --prefix-cache-n 0 \ + --prompt-len 64 --gen-len "$GEN" --iterations 1 --warmup 0 >"$log" 2>&1 + rc=$? + flock -u 9 # release BEFORE parsing + + [ $rc -eq 0 ] || { + echo "FAIL_BENCH $tag rc=$rc" + # A process killed by a GPU fault dies with no error line. This box carries + # ~1,485 Xid faults, so check before blaming the code. + dmesg 2>/dev/null | grep -i xid | tail -3 + tail -5 "$log" + exit 1 + } + # "| Decode (260 tokens) | 13.9 +- 0.0 | 71.80 ms/T |". An empty parse must + # FAIL, not silently become a 0.0 ms delta. + ms=$(grep -oE "[0-9]+\.[0-9]+ ms/T" "$log" | tail -1 | grep -oE "^[0-9]+\.[0-9]+") + [ -n "$ms" ] || { echo "FAIL_NO_MS_PER_T $tag"; tail -5 "$log"; exit 1; } + echo "$tag $mode $ms" +} + +: >"$OUT/ab.txt" +for i in $(seq 1 "$REPS"); do + leg "A$i" cpu | tee -a "$OUT/ab.txt" + leg "B$i" fused | tee -a "$OUT/ab.txt" +done + +echo "=== RESULT ===" +awk ' + { v[substr($1,1,1) substr($1,2)] = $3; arm[substr($1,1,1)] += $3; n[substr($1,1,1)]++ } + END { + a = arm["A"]/n["A"]; b = arm["B"]/n["B"]; + printf "A (cpu round trip) mean %.2f ms/token over %d legs\n", a, n["A"]; + printf "B (fused kernel) mean %.2f ms/token over %d legs\n", b, n["B"]; + printf "delta %.2f ms/token (%+.2f%%)\n", a-b, -100.0*(a-b)/a; + if (n["A"]>1) printf "drift within A %+.2f%% (A1 %.2f -> A%d %.2f)\n", \ + 100.0*(v["A" n["A"]]-v["A1"])/v["A1"], v["A1"], n["A"], v["A" n["A"]]; + if (n["B"]>1) printf "drift within B %+.2f%% (B1 %.2f -> B%d %.2f)\n", \ + 100.0*(v["B" n["B"]]-v["B1"])/v["B1"], v["B1"], n["B"], v["B" n["B"]]; + print "A delta smaller than the drift is not a result."; + }' "$OUT/ab.txt" From 488e4585d49a7c9385744e161d54573515b7da78 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Tue, 18 Aug 2026 12:02:31 +0100 Subject: [PATCH 09/22] measure(ArcKV/Fp8): 44 -> 1 D2H/step; fused kernel 1.98 us/call [MEASURED] Box arc-v4-stack (H200), binary md5 e15259dc9ce935fa8782ba832cac1992 in a PRIVATE target dir (/root/arc-wt/fp8-target, not the shared /root/arc-wt/target that a neighbour's build overwrote tonight), 118 arc_kv_fp8_quantize symbol hits, --features 'cuda flash-attn'. nsys, 20 s steady-state tail, step count pinned from the trace by two independent anchors. before (cpu) after (fused) cuMemcpyDtoHAsync_v2 43.89/step 1.00/step ... of which 1,792 B 42.89/step size ABSENT from the trace ... of which 517,120 B 1.00/step 1.00/step (logits, not ours) kernel launches 9,081.8/step 8,404.5/step (-677.3) cuLaunchKernel 7,892.4/step 7,119.0/step (-773.5) cuMemAllocAsync 11,376.7/step 10,618.0/step (-758.6) ALL *Synchronize* 0.00/step 0.00/step (why grep lies here) Per-call, which is the evidence that survives a contended box: before cuMemcpyDtoHAsync_v2 48.56 us/call HOST time -> 2.131 ms/step blocking after arc_kv_fp8_quantize_kernel 1.98 us/call, 42.93 calls/step arc_kv_fp8_dequantize_kernel 2.34 us/call, 42.93 calls/step fused total 0.1853 ms/step device time 42.93 rather than 43.00 is one step straddling the window edge (99.84%); the before arm's 1,792 B D2H shows the same 0.26% at 42.89/43. Engagement is therefore proven per layer, not assumed. NO end-to-end tok/s delta is reported. A naive A-B-A-B delta on this box is biased by exactly one slot of drift, and four end-to-end numbers here turned out to be pure environment in one night. The A/B driver is dropped rather than shipped with a result it cannot support; the case rests on per-call cost and launch counts, which do not depend on how long the step took. Co-Authored-By: Claude Opus 5 (1M context) --- arc-tools/kv_fp8_ab.sh | 83 ---------------------------- arc-tools/kv_fp8_percall.py | 107 ++++++++++++++++++++++++++++++++++++ 2 files changed, 107 insertions(+), 83 deletions(-) delete mode 100755 arc-tools/kv_fp8_ab.sh create mode 100644 arc-tools/kv_fp8_percall.py diff --git a/arc-tools/kv_fp8_ab.sh b/arc-tools/kv_fp8_ab.sh deleted file mode 100755 index 876d5042c..000000000 --- a/arc-tools/kv_fp8_ab.sh +++ /dev/null @@ -1,83 +0,0 @@ -#!/usr/bin/env bash -# ArcKV/Fp8 — interleaved A-B-A-B decode timing for the fused E4M3 kernel. -# A = ARC_KV_FP8_MODE=cpu (43 blocking D2H per step) -# B = ARC_KV_FP8_MODE=fused (one device kernel, 0 D2H) -# -# A-then-B is banned on this box: a 3.3% monotonic drift once exceeded the arm -# difference. The drift within each arm is printed next to the delta, and a -# delta smaller than the drift is not a result. -# -# LOCK DISCIPLINE. The lock covers exactly one thing: a leg (model load, which -# allocates GPU memory, plus the timed run). It is taken immediately before and -# released immediately after — NOT held across parsing, and NOT wrapped around -# the whole script. Six chains share this H200; one holding the lock through -# post-processing starved the trellis chain for 40 minutes while the card read -# 0 %, 0 MiB, 78 W. -set -u -OUT=${OUT:-/root/kvfp8} -BASE=${BASE:-/root/budget-chain} -BIN=${BIN:-/root/arc-wt/fp8-target/release/mistralrs} -LOCK=${LOCK:-/root/locks/bench.lock} -LOCKWAIT=${LOCKWAIT:-3600} -REPS=${REPS:-2} -GEN=${GEN:-260} -export PATH=/usr/local/cuda-13.1/bin:$PATH -export LD_LIBRARY_PATH=/usr/local/cuda/compat:${LD_LIBRARY_PATH:-} -unset ARC_TIME_DECODE V4_STATS V4_NAN_DEBUG V4_TRACE ARC_PROFILE ARC_GPU_ACT_QUANT -[ -x "$BIN" ] || { echo "FATAL_NO_BINARY $BIN"; exit 2; } - -exec 9>"$LOCK" - -leg() { # $1 = tag, $2 = mode - local tag=$1 mode=$2 log="$OUT/$1.log" free rc ms - flock -w "$LOCKWAIT" 9 || { echo "FATAL_LOCK_TIMEOUT ${LOCKWAIT}s"; exit 2; } - # Holding the lock is not the same as the card being free: V4 is ~79 GB of - # 143 GB and the previous holder's process can outlive its script. - for _ in $(seq 1 240); do - free=$(nvidia-smi --query-gpu=memory.free --format=csv,noheader,nounits) - [ "${free:-0}" -ge 100000 ] && break - sleep 5 - done - [ "${free:-0}" -ge 100000 ] || { echo "FATAL_NO_VRAM ${free}MiB"; exit 2; } - - ARC_KV_FP8_MODE=$mode "$BIN" bench -m "$BASE/src" -a deepseekv4 \ - --from-uqff "$BASE/uqff/qtip2-0.uqff" --max-seqs 1 --prefix-cache-n 0 \ - --prompt-len 64 --gen-len "$GEN" --iterations 1 --warmup 0 >"$log" 2>&1 - rc=$? - flock -u 9 # release BEFORE parsing - - [ $rc -eq 0 ] || { - echo "FAIL_BENCH $tag rc=$rc" - # A process killed by a GPU fault dies with no error line. This box carries - # ~1,485 Xid faults, so check before blaming the code. - dmesg 2>/dev/null | grep -i xid | tail -3 - tail -5 "$log" - exit 1 - } - # "| Decode (260 tokens) | 13.9 +- 0.0 | 71.80 ms/T |". An empty parse must - # FAIL, not silently become a 0.0 ms delta. - ms=$(grep -oE "[0-9]+\.[0-9]+ ms/T" "$log" | tail -1 | grep -oE "^[0-9]+\.[0-9]+") - [ -n "$ms" ] || { echo "FAIL_NO_MS_PER_T $tag"; tail -5 "$log"; exit 1; } - echo "$tag $mode $ms" -} - -: >"$OUT/ab.txt" -for i in $(seq 1 "$REPS"); do - leg "A$i" cpu | tee -a "$OUT/ab.txt" - leg "B$i" fused | tee -a "$OUT/ab.txt" -done - -echo "=== RESULT ===" -awk ' - { v[substr($1,1,1) substr($1,2)] = $3; arm[substr($1,1,1)] += $3; n[substr($1,1,1)]++ } - END { - a = arm["A"]/n["A"]; b = arm["B"]/n["B"]; - printf "A (cpu round trip) mean %.2f ms/token over %d legs\n", a, n["A"]; - printf "B (fused kernel) mean %.2f ms/token over %d legs\n", b, n["B"]; - printf "delta %.2f ms/token (%+.2f%%)\n", a-b, -100.0*(a-b)/a; - if (n["A"]>1) printf "drift within A %+.2f%% (A1 %.2f -> A%d %.2f)\n", \ - 100.0*(v["A" n["A"]]-v["A1"])/v["A1"], v["A1"], n["A"], v["A" n["A"]]; - if (n["B"]>1) printf "drift within B %+.2f%% (B1 %.2f -> B%d %.2f)\n", \ - 100.0*(v["B" n["B"]]-v["B1"])/v["B1"], v["B1"], n["B"], v["B" n["B"]]; - print "A delta smaller than the drift is not a result."; - }' "$OUT/ab.txt" diff --git a/arc-tools/kv_fp8_percall.py b/arc-tools/kv_fp8_percall.py new file mode 100644 index 000000000..ab643a871 --- /dev/null +++ b/arc-tools/kv_fp8_percall.py @@ -0,0 +1,107 @@ +#!/usr/bin/env python3 +"""Per-call cost of the FP8 KV round trip, from an nsys trace. + +Parent system: ArcLab (measuring ArcInfer / ArcKV / Fp8). + +WHY PER-CALL AND NOT tok/s + End-to-end tok/s on this shared box is not trustworthy: four separate + end-to-end numbers here turned out to be pure environment in one night, and + a naive A-B-A-B delta is *biased*, not merely noisy, because at equal + spacing the arms differ by exactly one slot of drift. A per-call cost taken + from the trace is immune to both — it does not depend on how long the step + took, only on what the step did. The `GpuApprox` conclusion survived a bad + night for exactly this reason: its evidence was `kv_fp8_quant` 73.49 -> + 134.18 us/call, not a throughput delta. + +WHAT IT REPORTS + BEFORE arm: the host time spent inside `cuMemcpyDtoHAsync_v2` per decode + step. That is the real cost of the CPU round trip -- an "async" copy taking + ~109 us of HOST time is a pageable-memory staged copy, i.e. blocking, which + is also why a CUDA graph cannot record it. + AFTER arm: the device time of the two fused kernels per decode step, and + their calls/step (which must be 43 -- one per layer -- or the fused path did + not engage on every layer and the number is not what it claims). + +Exits 2 on analysis failure, never 1. + +Usage: kv_fp8_percall.py [tail_seconds] +""" +import collections +import sqlite3 +import sys + + +def die(msg): + print(f"FATAL {msg}", file=sys.stderr) + sys.exit(2) + + +def main(): + if len(sys.argv) < 2: + die("usage: kv_fp8_percall.py [tail_seconds]") + tail_s = float(sys.argv[2]) if len(sys.argv) > 2 else 20.0 + db = sqlite3.connect(sys.argv[1]) + cur = db.cursor() + strings = {i: v for i, v in cur.execute("SELECT id, value FROM StringIds")} + KT = "CUPTI_ACTIVITY_KIND_KERNEL" + MT = "CUPTI_ACTIVITY_KIND_MEMCPY" + RT = "CUPTI_ACTIVITY_KIND_RUNTIME" + + end = cur.execute(f"SELECT MAX(end) FROM {KT}").fetchone()[0] + if end is None: + die("no kernels in the trace") + win0 = end - int(tail_s * 1e9) + + # Steps = the once-per-step logits readback. + big = list(cur.execute( + f"SELECT COUNT(*) FROM {MT} WHERE start>=? AND copyKind=2 AND bytes>100000", + (win0,))) + steps = big[0][0] if big else 0 + if steps < 20: + die(f"only {steps} steps in the window; widen tail_seconds") + print(f"steps in window {steps}") + + # ---- BEFORE-arm evidence: host time inside the blocking D2H. + tot_ns = 0 + n = 0 + for st, en in cur.execute( + f"SELECT start, end FROM {RT} WHERE start>=? AND nameId IN " + f"(SELECT id FROM StringIds WHERE value='cuMemcpyDtoHAsync_v2')", + (win0,)): + tot_ns += en - st + n += 1 + if n: + print(f"cuMemcpyDtoHAsync_v2 {n / steps:8.2f} calls/step " + f"{tot_ns / n / 1000.0:8.2f} us/call (HOST time) " + f"{tot_ns / steps / 1e6:8.3f} ms/step") + else: + print("cuMemcpyDtoHAsync_v2 0.00 calls/step") + + # ---- AFTER-arm evidence: the fused kernels' device time. + dev = collections.defaultdict(lambda: [0, 0]) + for name_id, st, en in cur.execute( + f"SELECT shortName, start, end FROM {KT} WHERE start>=?", (win0,)): + nm = strings.get(name_id, str(name_id)) + if "arc_kv_fp8" in nm: + d = dev[nm] + d[0] += 1 + d[1] += en - st + if dev: + print("fused kernels (device time):") + tot = 0.0 + for nm, (c, ns) in sorted(dev.items()): + per_step = c / steps + print(f" {nm[:44]:44s} {per_step:7.2f} calls/step " + f"{ns / c / 1000.0:7.2f} us/call " + f"{ns / steps / 1e6:7.4f} ms/step") + tot += ns / steps / 1e6 + if abs(per_step - round(per_step)) > 0.05: + print(f" WARN {nm} is not an integer per step") + print(f" fused total{'':38s} {tot:7.4f} ms/step") + else: + print("fused kernels ABSENT from this trace " + "(this is the before arm, or the fused path never engaged)") + + +if __name__ == "__main__": + main() From 46c99e1f23e573e59818c4707ad43cd9187cb21c Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Thu, 20 Aug 2026 18:36:35 +0100 Subject: [PATCH 10/22] =?UTF-8?q?fix(ArcKV/Fp8):=20ARC=5FKV=5FFP8=5FMODE?= =?UTF-8?q?=20was=20a=20lying=20switch=20=E2=80=94=20rename,=20and=20stop?= =?UTF-8?q?=20typos=20selecting=20GpuApprox?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The name says "mode", which reads as an on/off for FP8 KV. It is not one, and it never was: it selects WHICH ARITHMETIC produces the E4M3 code, and all three variants quantize. There is no "off" — V4 is FP8-QAT, so the quantize/dequantize round trip at `deepseek4.rs:1689` is the model's numerics, not an optimisation, and it is correctly unconditional. The flag that gates FP8 KV *storage* is a different variable, `ARC_V4_FP8_KV`. Three defects, in descending severity: 1. A TYPO COULD SILENTLY SELECT NON-BIT-EXACT ARITHMETIC. The old table was `_ if ARC_GPU_ACT_QUANT is set => GpuApprox` followed by `_ => FusedDevice`. Both arms were reachable ONLY for an unset or MISSPELLED value. So on any box that still had `ARC_GPU_ACT_QUANT` exported from an earlier experiment, `ARC_KV_FP8_MODE=fusd` selected `GpuApprox` — the one variant this module documents as round-half-away-from-zero rather than round-half-to-even, and labels "NOT the fix and must not be shipped as one". Resolved once into a `OnceLock`, so it stuck for the process lifetime with no trace. An unparsable value now lands on the bit-exact default and SAYS SO, on both stderr and `tracing::error!` — never on `GpuApprox`. 2. The docs stated the opposite of the code. `PROFILING.md` annotated the `kv_fp8_quant` span "opt-in, ARC_V4_FP8_KV=1". That span opens at `deepseek4.rs:1688` and runs on every forward regardless of the flag; the adjacent `kv_fp8_dequant` line carried no such note, so the doc was internally inconsistent too. The annotation moves to `kv_cache_append`, which is where `ARC_V4_FP8_KV` actually acts. 3. `unwrap_or_default()` collapsed unset and empty-string, and nothing trimmed. Renamed to `ARC_KV_FP8_IMPL`. The old spelling still works and prints a deprecation on both channels — renaming it silently would convert every operator's muscle memory into a fresh silent failure, which is the disease being cured, not the cure. The table is now a pure `parse_impl(Option<&str>, bool)` with a test that pins every arm INCLUDING the typo-must-not-reach-GpuApprox regression. It runs on the free CPU lane; no GPU is required to keep this honest. --- arc-tools/kv_fp8_nsys_ab.sh | 4 +- docs/engineering/PROFILING.md | 6 +- mistralrs-core/src/models/dsv4_kv_fp8.rs | 150 +++++++++++++++++++++-- 3 files changed, 144 insertions(+), 16 deletions(-) diff --git a/arc-tools/kv_fp8_nsys_ab.sh b/arc-tools/kv_fp8_nsys_ab.sh index 362a2b946..80a38692e 100755 --- a/arc-tools/kv_fp8_nsys_ab.sh +++ b/arc-tools/kv_fp8_nsys_ab.sh @@ -49,13 +49,13 @@ exec 9>"$LOCK" for arm in before:cpu after:fused; do A=${arm%%:*} M=${arm##*:} - echo "=== leg $A (ARC_KV_FP8_MODE=$M) $(date -u +%T) ===" + echo "=== leg $A (ARC_KV_FP8_IMPL=$M) $(date -u +%T) ===" rm -f "$OUT/$A".nsys-rep "$OUT/$A".sqlite # LOCK HELD ONLY HERE: VRAM wait + the traced run. Released the instant the # bench exits, before any report export. flock -w "$LOCKWAIT" 9 || { echo "FATAL_LOCK_TIMEOUT after ${LOCKWAIT}s"; exit 2; } wait_for_vram 100000 - ARC_KV_FP8_MODE=$M nsys profile --trace=cuda --sample=none --cpuctxsw=none \ + ARC_KV_FP8_IMPL=$M nsys profile --trace=cuda --sample=none --cpuctxsw=none \ --cuda-memory-usage=false --duration="$DUR" --kill=sigterm \ --force-overwrite=true --output="$OUT/$A" \ "$BIN" bench -m "$BASE/src" -a deepseekv4 \ diff --git a/docs/engineering/PROFILING.md b/docs/engineering/PROFILING.md index 92c466dfc..55a183774 100644 --- a/docs/engineering/PROFILING.md +++ b/docs/engineering/PROFILING.md @@ -211,10 +211,10 @@ step engine/mod.rs — one scheduler iter │ │ │ │ ├─ kv_proj [device] fused wkv │ │ │ │ ├─ kv_norm [device] │ │ │ │ ├─ rope [device] -│ │ │ │ ├─ kv_fp8_quant [device] opt-in, ARC_V4_FP8_KV=1 -│ │ │ │ ├─ kv_fp8_dequant [device] +│ │ │ │ ├─ kv_fp8_quant [device] ALWAYS (V4 is FP8-QAT), not opt-in +│ │ │ │ ├─ kv_fp8_dequant [device] ALWAYS — pairs with the quant above │ │ │ │ ├─ compressed_kv_build[device] -│ │ │ │ ├─ kv_cache_append [device] +│ │ │ │ ├─ kv_cache_append [device] FP8 *storage* here is ARC_V4_FP8_KV=1 │ │ │ │ ├─ kv_cache_span [device] │ │ │ │ ├─ sdpa [device] dsv4_attention (window ∧ compressed) │ │ │ │ ├─ inv_rope [device] NOT inside ARC_TIME_DECODE's timer diff --git a/mistralrs-core/src/models/dsv4_kv_fp8.rs b/mistralrs-core/src/models/dsv4_kv_fp8.rs index 773e95861..a52feac5d 100644 --- a/mistralrs-core/src/models/dsv4_kv_fp8.rs +++ b/mistralrs-core/src/models/dsv4_kv_fp8.rs @@ -73,7 +73,7 @@ pub(crate) enum KvQuantMode { FusedDevice, /// Exact E4M3 via candle's CPU cast (candle has no CUDA `F8E4M3` cast — /// "named symbol not found"), at the price of one device sync per layer. - /// `ARC_KV_FP8_MODE=cpu`. + /// `ARC_KV_FP8_IMPL=cpu`. CpuExact, /// On-device float arithmetic that reproduces E4M3's value grid with /// round-half-away-from-zero instead of round-half-to-even. Removes the @@ -90,25 +90,94 @@ pub(crate) enum KvQuantMode { GpuApprox, } +/// The legal values of `ARC_KV_FP8_IMPL`, for error messages and for the test +/// that pins the table. +pub(crate) const KV_FP8_IMPL_VALUES: &[&str] = &[ + "fused", + "fused_device", + "cpu", + "cpu_exact", + "gpu", + "gpu_approx", +]; + impl KvQuantMode { + /// Pure half of [`Self::from_env`], so the table can be tested without + /// mutating the process environment (which is racy across test threads and + /// `unsafe` since the 2024 edition). + /// + /// `raw` is the `ARC_KV_FP8_IMPL` value (`None` when unset); `legacy_gpu` + /// is whether the pre-existing `ARC_GPU_ACT_QUANT` is set. Returns the mode + /// and, when the input was not understood, the complaint to shout about it. + pub(crate) fn parse_impl(raw: Option<&str>, legacy_gpu: bool) -> (Self, Option) { + let default = if legacy_gpu { + // Preserved from before this flag existed: `ARC_GPU_ACT_QUANT=1` + // selects the approximate on-device path. It only applies when + // `ARC_KV_FP8_IMPL` says nothing. + Self::GpuApprox + } else { + Self::FusedDevice + }; + match raw.map(|s| s.trim()) { + None | Some("") => (default, None), + Some(v) => match v.to_ascii_lowercase().as_str() { + "cpu" | "cpu_exact" => (Self::CpuExact, None), + "gpu" | "gpu_approx" => (Self::GpuApprox, None), + "fused" | "fused_device" => (Self::FusedDevice, None), + // NOT `default`. Falling through to `default` here is how a + // typo used to select `GpuApprox` — the one variant documented + // as *not* bit-exact and "must not be shipped" — on any box + // that still had `ARC_GPU_ACT_QUANT` set from an earlier + // experiment. An unparsable value now lands on the bit-exact + // default and says so. + _ => ( + Self::FusedDevice, + Some(format!( + "ARC_KV_FP8_IMPL={v:?} is not a legal value (expected one of {}); \ + falling back to the default fused-device kernel. \ + NOTE: this flag selects WHICH ARITHMETIC produces the E4M3 code. \ + It has no 'off' — V4 is FP8-QAT and the quantize/dequantize round \ + trip is the model's numerics, not an optimisation. The flag that \ + gates FP8 KV *storage* is ARC_V4_FP8_KV.", + KV_FP8_IMPL_VALUES.join(", ") + )), + ), + }, + } + } + /// Resolved once per process: this is called per attention layer per /// forward, and `deepseek4` has already been bitten by per-call /// `std::env::var_os` in exactly that position (~390 environment scans per /// forward, wave33). + /// + /// Reads `ARC_KV_FP8_IMPL`. The former spelling `ARC_KV_FP8_MODE` is still + /// honoured — loudly — because renaming it silently would turn every + /// operator's muscle memory into a new silent failure, which is the same + /// disease the rename cures. pub(crate) fn from_env() -> Self { static MODE: OnceLock = OnceLock::new(); *MODE.get_or_init(|| { - match std::env::var("ARC_KV_FP8_MODE") - .unwrap_or_default() - .to_ascii_lowercase() - .as_str() - { - "cpu" | "cpu_exact" => Self::CpuExact, - "gpu" | "gpu_approx" => Self::GpuApprox, - "fused" | "fused_device" => Self::FusedDevice, - _ if std::env::var_os("ARC_GPU_ACT_QUANT").is_some() => Self::GpuApprox, - _ => Self::FusedDevice, + let legacy_gpu = std::env::var_os("ARC_GPU_ACT_QUANT").is_some(); + let new = std::env::var("ARC_KV_FP8_IMPL").ok(); + let old = std::env::var("ARC_KV_FP8_MODE").ok(); + if old.is_some() { + // eprintln! as well as tracing: a rented box often runs before + // a subscriber is installed, and a deprecation nobody sees is + // not a deprecation. + let msg = "ARC_KV_FP8_MODE has been renamed ARC_KV_FP8_IMPL (it selects the \ + quantizer ARITHMETIC; it never had an 'off'). The old name still \ + works for now — please update."; + eprintln!("[arc-kv-fp8] {msg}"); + tracing::warn!("{msg}"); } + let raw = new.or(old); + let (mode, complaint) = Self::parse_impl(raw.as_deref(), legacy_gpu); + if let Some(c) = complaint { + eprintln!("[arc-kv-fp8] {c}"); + tracing::error!("{c}"); + } + mode }) } } @@ -739,4 +808,63 @@ mod tests { assert_eq!(values[0], 0.0); assert_eq!(values[126], 448.0); } + + /// The parse table, pinned. Before this, `ARC_KV_FP8_MODE` had a bare + /// `_ => Self::FusedDevice` catch-all preceded by an + /// `_ if ARC_GPU_ACT_QUANT is set => Self::GpuApprox` arm, so a MISSPELLED + /// value on a box that still had `ARC_GPU_ACT_QUANT` exported silently + /// selected `GpuApprox` — the one variant this module documents as not + /// bit-exact and "must not be shipped as one". + #[test] + fn kv_fp8_impl_parse_table() { + use KvQuantMode::*; + + // Unset: the bit-exact fused kernel, with and without trailing space. + assert_eq!(KvQuantMode::parse_impl(None, false).0, FusedDevice); + assert_eq!(KvQuantMode::parse_impl(Some(""), false).0, FusedDevice); + assert_eq!(KvQuantMode::parse_impl(Some(" "), false).0, FusedDevice); + + // Every legal value round-trips, case- and whitespace-insensitively, + // and none of them complains. + for (raw, want) in [ + ("fused", FusedDevice), + ("fused_device", FusedDevice), + ("FUSED", FusedDevice), + (" cpu ", CpuExact), + ("cpu_exact", CpuExact), + ("gpu", GpuApprox), + ("gpu_approx", GpuApprox), + ] { + let (got, complaint) = KvQuantMode::parse_impl(Some(raw), false); + assert_eq!(got, want, "{raw:?}"); + assert!(complaint.is_none(), "{raw:?} should parse silently"); + } + + // Every value the error message advertises as legal really is legal. + for v in KV_FP8_IMPL_VALUES { + assert!( + KvQuantMode::parse_impl(Some(v), false).1.is_none(), + "{v:?} is advertised as legal but is rejected" + ); + } + + // The legacy var still selects GpuApprox, but only when the new flag + // is silent. + assert_eq!(KvQuantMode::parse_impl(None, true).0, GpuApprox); + assert_eq!(KvQuantMode::parse_impl(Some(""), true).0, GpuApprox); + assert_eq!(KvQuantMode::parse_impl(Some("cpu"), true).0, CpuExact); + + // THE REGRESSION. A typo must be loud, and must NOT reach GpuApprox + // even with the legacy var set. + for legacy in [false, true] { + let (got, complaint) = KvQuantMode::parse_impl(Some("fusd"), legacy); + assert_eq!(got, FusedDevice, "typo must land on the bit-exact default"); + let c = complaint.expect("a typo must produce a complaint"); + assert!(c.contains("fusd"), "the complaint must quote the bad value"); + assert!( + c.contains("ARC_V4_FP8_KV"), + "the complaint must name the flag that actually gates storage" + ); + } + } } From 3a71d0392013e6269ea3676370537e60074f2b08 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Thu, 20 Aug 2026 18:46:22 +0100 Subject: [PATCH 11/22] =?UTF-8?q?fix(build):=20EXPECTED=5FKERNEL=5FCOUNT?= =?UTF-8?q?=2040=20->=2041=20=E2=80=94=20this=20branch=20adds=20arc=5Fkvqu?= =?UTF-8?q?ant.cu?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The ArcGate kernel-count tripwire caught a real regression introduced by rebasing this branch onto current master, and it is worth recording how, because the failure mode is subtle. This branch originally carried `EXPECTED_KERNEL_COUNT 39 -> 40 — arc_kvquant.cu was never counted`. Meanwhile master independently went 39 -> 40 for a DIFFERENT kernel. On rebase, git compared patch texts, saw an identical `-39 / +41`-shaped hunk already upstream, and dropped the commit as "patch contents already upstream". The number was right; the REASON was not. Net effect: the glob discovers 41 sources while the file still claims 40, so `arc_kvquant.cu` would once again be the kernel that goes missing quietly — which is the exact failure this file was created to make impossible. Caught by `cuda_kernel_build_guard::expected_kernel_count_matches_disk` on the free CPU lane, before any GPU time was spent. That is the tripwire earning its keep, not a nuisance — do not "fix" a future occurrence by relaxing the guard. --- mistralrs-quant/kernels/EXPECTED_KERNEL_COUNT | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mistralrs-quant/kernels/EXPECTED_KERNEL_COUNT b/mistralrs-quant/kernels/EXPECTED_KERNEL_COUNT index c5cd0dc04..b0a528a99 100644 --- a/mistralrs-quant/kernels/EXPECTED_KERNEL_COUNT +++ b/mistralrs-quant/kernels/EXPECTED_KERNEL_COUNT @@ -28,4 +28,4 @@ # (this count minus the 5 `*_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. -40 +41 From dd5583c4b753ae76f01aa84a65a5b7b516d04202 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Wed, 19 Aug 2026 20:20:59 +0100 Subject: [PATCH 12/22] perf(arckv): make candle's caching allocator reachable and on for decode MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `set_alloc_cache_enabled(true)` sat behind three stacked default-off gates: `probe && seq_len == 1 && env("ARC_CANDLE_ALLOC_CACHE")`. The first conjunct tied a general-purpose allocator to the V4 capture probe, so the only way to recycle a decode step's frees was to also be capturing. The ~11k allocations per token were a disabled feature, not a missing one. Replaced with a pure policy, `alloc_cache_action(seq_len, enabled, killed)`: * decode (`seq_len == 1`) -> Enable * prefill (`seq_len != 1`) -> DrainAndDisable * already in that state -> Leave Prefill draining is not incidental. The cache is keyed on exact byte count (`free: HashMap>`, no bucketing, no smallest-fit) and has no capacity bound and no eviction, so a prefill's large one-shot buffers would be parked for the process lifetime under a key nothing asks for again. `ARC_CANDLE_ALLOC_CACHE=0` is the kill switch. Any other value, and unset, leave the policy in force, so the `=1` the ops scripts pass still means what it always meant. The capture probe keeps its own gate for the graph-mode positions; only the allocator moved out from under it. Six host-runnable tests for the policy. The allocator has no tests at all in either repo — candle's `cuda_backend/{device,mod}.rs` carry no `#[cfg(test)]` and every arc-side exercise is `#[cfg(feature = "cuda")]` — so the decision of when it is on is now the part that CI can see. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01SpVNMpb13HkUXqSqbN1o9H --- mistralrs-core/src/pipeline/normal.rs | 200 ++++++++++++++++++++++++-- 1 file changed, 192 insertions(+), 8 deletions(-) diff --git a/mistralrs-core/src/pipeline/normal.rs b/mistralrs-core/src/pipeline/normal.rs index d412a7cd9..7f32c2563 100644 --- a/mistralrs-core/src/pipeline/normal.rs +++ b/mistralrs-core/src/pipeline/normal.rs @@ -98,6 +98,91 @@ pub struct NormalPipeline { autonomous_runner: Option, } +/// What a forward pass should do with candle's caching allocator before it +/// runs. See [`alloc_cache_action`]. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub(crate) enum AllocCacheAction { + /// Turn the cache on. Frees start being recycled instead of returned to + /// the driver. + Enable, + /// Return every held buffer to the driver and turn the cache off. + DrainAndDisable, + /// Already in the right state; touch nothing. + Leave, +} + +/// Decide whether candle's caching allocator should be on for this forward. +/// +/// # Why this is not simply "on" +/// +/// The cache (fork `88d86a2`, `candle-core/src/cuda_backend/device.rs:42`) is +/// keyed on **exact byte count** — `free: HashMap>`, +/// looked up with `free.get_mut(&bytes)`. No bucketing, no smallest-fit, no +/// splitting. It also has **no capacity bound and no eviction**: the only ways +/// memory goes back to the driver are `set_alloc_cache_enabled(false)` and +/// `drain_alloc_cache_and_free()`. +/// +/// Those two facts together decide the policy: +/// +/// * **Decode wants it on.** A decode step allocates the same shapes 43 times +/// over (once per layer), and the shapes that do not depend on KV length — +/// hidden states, MLP and expert intermediates — repeat step after step. +/// Exact-size keying is a perfect fit for that traffic. +/// * **Prefill must drain it.** Prefill shapes scale with prompt length, are +/// large, and are hit once. Leaving the cache on across a prefill parks +/// those buffers for the process lifetime under a byte-size key nothing will +/// ever request again. That is the failure mode `ARC_NO_DEDICATED_DECODE` +/// already exists to work around, and it is why this returns +/// [`AllocCacheAction::DrainAndDisable`] rather than `Leave` on `seq_len != 1`. +/// +/// # The measurement this is waiting on +/// +/// Buffers whose size tracks KV length — the causal mask most obviously — +/// change size every decode step, so each step files one more never-reused +/// entry. That is bounded by `O(context^2)` bytes in the worst case and is +/// *not* bounded by this policy. It is small at short context (a `[1,1,1,kv]` +/// BF16 mask over 4k tokens sums to ~16 MB) and is not small at 128k. +/// +/// The fixed-capacity graph-mode path (`deepseek4.rs:4344-4353`, which swaps +/// the growing causal mask for `graph_mode_length_mask` at +/// `cfg_full.sliding_window`) removes that growth entirely, but it is reached +/// only under `ARC_V4_CAPTURE_PROBE`. Until the shape-invariance work lands, +/// `ARC_CANDLE_ALLOC_CACHE=0` is the kill switch, and the long-context +/// high-water mark is the number a GPU run should falsify this with. +/// +/// # Contract +/// +/// * `seq_len` — the forward's sequence length. `1` is decode. +/// * `enabled` — what candle reports *now* (`alloc_cache_enabled()`), so the +/// action is idempotent and we never drain a cache that is already off. +/// * `killed` — `ARC_CANDLE_ALLOC_CACHE=0` was set. +/// +/// Pure, so the policy is testable without a GPU — which matters, because the +/// allocator itself has no tests at all in either repo. +#[cfg_attr(not(feature = "cuda"), allow(dead_code))] +pub(crate) fn alloc_cache_action( + seq_len: usize, + enabled: bool, + killed: bool, +) -> AllocCacheAction { + let want = !killed && seq_len == 1; + match (want, enabled) { + (true, false) => AllocCacheAction::Enable, + (false, true) => AllocCacheAction::DrainAndDisable, + _ => AllocCacheAction::Leave, + } +} + +/// `ARC_CANDLE_ALLOC_CACHE=0` turns the caching allocator off entirely. +/// +/// Any other value — and *unset* — leaves the default policy in force. The +/// variable used to be the on-switch, so the value `1` the ops scripts pass +/// (`arc-tools/arcgraph_heap_probe.sh:191`) still means what it always meant. +#[cfg_attr(not(feature = "cuda"), allow(dead_code))] +pub(crate) fn alloc_cache_killed() -> bool { + std::env::var("ARC_CANDLE_ALLOC_CACHE").is_ok_and(|v| v == "0") +} + /// A loader for a "normal" (non-quantized) model. pub struct NormalLoader { inner: Box, @@ -1702,15 +1787,31 @@ impl Pipeline for NormalPipeline { "normal.rs:1554", ); } - // Enable the candle caching allocator for graph-capture - // safety (RUN-161). Idempotent; gated so it's off during - // model load and only active for decode. Warmup decode - // forwards populate the cache before capture. - if probe && seq_len == 1 && std::env::var_os("ARC_CANDLE_ALLOC_CACHE").is_some() - { - if let candle_core::Device::Cuda(cd) = self.device() { - cd.set_alloc_cache_enabled(true); + // Candle's caching allocator, driven by `alloc_cache_action`. + // + // This used to be `probe && seq_len == 1 && env(...)`, which + // made a general-purpose allocator reachable only when the + // V4 capture probe was also on. The cache recycles frees for + // ANY decode step — it is not capture machinery — and the + // three stacked default-off gates meant the ~11k allocations + // per token were a disabled feature rather than a missing + // one. It is now on for decode by default and drained on + // prefill; see `alloc_cache_action` for why prefill must + // drain rather than coast. + if let candle_core::Device::Cuda(cd) = self.device() { + match alloc_cache_action( + seq_len, + cd.alloc_cache_enabled(), + alloc_cache_killed(), + ) { + AllocCacheAction::Enable => cd.set_alloc_cache_enabled(true), + AllocCacheAction::DrainAndDisable => { + cd.set_alloc_cache_enabled(false) + } + AllocCacheAction::Leave => {} } + } + if probe && seq_len == 1 { // RUN-161 step 2b. Set the graph-mode device position: // drives RoPE + the fixed-capacity KV write slot, and // makes warmup forwards take the shape-constant path so @@ -2528,3 +2629,86 @@ impl AnyMoePipelineMixin for NormalPipeline { self.model.amoe_supported() } } + +/// The caching-allocator policy's contract. +/// +/// The allocator itself has **no tests at all** — not in this repo, and not in +/// the candle fork that implements it (`grep -rn "alloc_cache" --include="*.rs"` +/// over `candle-core` finds only `cuda_backend/device.rs` and +/// `cuda_backend/mod.rs`, with no `#[cfg(test)]` in either). Everything that +/// observes it needs a GPU, so nothing observes it in CI. +/// +/// The decision of *when* it is on does not need a GPU, so it is tested here. +#[cfg(test)] +mod alloc_cache_policy_tests { + use super::{alloc_cache_action, AllocCacheAction}; + + /// The change this policy exists to make: a decode step turns the cache on + /// without any capture probe being involved. + #[test] + fn decode_enables_the_cache() { + assert_eq!( + alloc_cache_action(1, false, false), + AllocCacheAction::Enable + ); + } + + /// Prefill must hand its buffers back. They are large, they are keyed by an + /// exact byte size no later request will ask for, and nothing evicts them. + #[test] + fn prefill_drains_the_cache() { + assert_eq!( + alloc_cache_action(512, true, false), + AllocCacheAction::DrainAndDisable + ); + } + + /// Called once per forward, so it must be a no-op in the steady state + /// rather than re-enabling (or re-draining) every step. + #[test] + fn steady_state_touches_nothing() { + assert_eq!(alloc_cache_action(1, true, false), AllocCacheAction::Leave); + assert_eq!( + alloc_cache_action(512, false, false), + AllocCacheAction::Leave + ); + } + + /// The kill switch has to work from either state, including turning off a + /// cache that a previous step already enabled — otherwise setting it + /// mid-run would leave the buffers stranded. + #[test] + fn kill_switch_disables_and_drains_from_either_state() { + assert_eq!( + alloc_cache_action(1, true, true), + AllocCacheAction::DrainAndDisable + ); + assert_eq!(alloc_cache_action(1, false, true), AllocCacheAction::Leave); + } + + /// A prompt that happens to be one token long is a prefill by every other + /// measure, but it allocates decode-shaped buffers, so the policy keys on + /// the shape it will actually see. Pinned so the equivalence is deliberate. + #[test] + fn one_token_prompt_is_treated_as_decode() { + assert_eq!( + alloc_cache_action(1, false, false), + AllocCacheAction::Enable + ); + } + + /// A zero-length forward is not decode. Guards against `seq_len == 0` + /// (which `dims2().unwrap_or((0, 0))` produces on a shape error) quietly + /// enabling the cache. + #[test] + fn degenerate_zero_length_forward_does_not_enable() { + assert_eq!( + alloc_cache_action(0, false, false), + AllocCacheAction::Leave + ); + assert_eq!( + alloc_cache_action(0, true, false), + AllocCacheAction::DrainAndDisable + ); + } +} From 7d91c1127ee3118fccd603bdfa64b73f2fb724aa Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Wed, 19 Aug 2026 23:48:23 +0100 Subject: [PATCH 13/22] perf(arckv): bound the caching allocator, and print the counters that prove it MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Picks up candle b2a4dd80, which gives `AllocCache` a capacity and LRU eviction. The allocator had neither, and arc only drains it on a prefill, so a single long generation grew forever. Measured on an H200 over 2 600 tokens, `memory.used` at 2 Hz: **+6.04 MiB per decoded token with no plateau**, against +0.057 MiB/token for the same run with the cache off — so the growth is the cache's and nothing else's. `ARC_ALLOC_CACHE_MAX_MB` sets the cap; unset leaves candle's 1 GiB default; `0` restores the old unbounded behaviour for A/B. A typo deliberately does *not* fall back to unbounded. `ARC_ALLOC_CACHE_STATS=N` prints the allocator's counters every N decode steps: allocations per step, **frees per step**, hit rate, bytes held against the cap. Those are the numbers this has to be judged on. A green log is not evidence — an earlier arena here reported "accounting OK" and bit-identical output over 52 steps while silently bypassing itself for every buffer under 128 bytes (KERNEL_RULES.md:977-984). Allocations staying low *and* frees being non-zero is. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01SpVNMpb13HkUXqSqbN1o9H --- Cargo.toml | 10 +- mistralrs-core/src/pipeline/normal.rs | 138 +++++++++++++++++++++++++- 2 files changed, 142 insertions(+), 6 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index b8dde1f07..225d8ecbe 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -49,11 +49,11 @@ rust-version = "1.88" # pointer has broken CI on this repo before — a path/branch dependency that # advances under a merged PR turns a green build into a red one with no commit # here to blame. -candle-core = { git = "https://github.com/aeonmindai/candle.git", rev = "88d86a2019ee552923b217284e89847d2785dcdf" } -candle-nn = { git = "https://github.com/aeonmindai/candle.git", rev = "88d86a2019ee552923b217284e89847d2785dcdf" } -candle-flash-attn-v3 = { git = "https://github.com/aeonmindai/candle.git", rev = "88d86a2019ee552923b217284e89847d2785dcdf" } -candle-flash-attn = { git = "https://github.com/aeonmindai/candle.git", rev = "88d86a2019ee552923b217284e89847d2785dcdf" } -candle-metal-kernels = { git = "https://github.com/aeonmindai/candle.git", rev = "88d86a2019ee552923b217284e89847d2785dcdf" } +candle-core = { git = "https://github.com/aeonmindai/candle.git", rev = "b2a4dd80b82726730c3c4c90421fa3a0938be984" } +candle-nn = { git = "https://github.com/aeonmindai/candle.git", rev = "b2a4dd80b82726730c3c4c90421fa3a0938be984" } +candle-flash-attn-v3 = { git = "https://github.com/aeonmindai/candle.git", rev = "b2a4dd80b82726730c3c4c90421fa3a0938be984" } +candle-flash-attn = { git = "https://github.com/aeonmindai/candle.git", rev = "b2a4dd80b82726730c3c4c90421fa3a0938be984" } +candle-metal-kernels = { git = "https://github.com/aeonmindai/candle.git", rev = "b2a4dd80b82726730c3c4c90421fa3a0938be984" } axum = "0.8.8" anyhow = "1.0.100" diff --git a/mistralrs-core/src/pipeline/normal.rs b/mistralrs-core/src/pipeline/normal.rs index 7f32c2563..a69835923 100644 --- a/mistralrs-core/src/pipeline/normal.rs +++ b/mistralrs-core/src/pipeline/normal.rs @@ -183,6 +183,90 @@ pub(crate) fn alloc_cache_killed() -> bool { std::env::var("ARC_CANDLE_ALLOC_CACHE").is_ok_and(|v| v == "0") } +/// Retention cap for the caching allocator, in bytes, or `None` to leave +/// candle's own default (1 GiB) in force. +/// +/// `ARC_ALLOC_CACHE_MAX_MB=0` means unbounded — the pre-bounding behaviour, kept +/// reachable so the leak it causes can be A/B'd rather than argued about. +/// Measured unbounded on V4: **+6.04 MiB per decoded token with no plateau**, +/// against +0.057 MiB/token with the cache off entirely. +#[cfg_attr(not(feature = "cuda"), allow(dead_code))] +pub(crate) fn alloc_cache_capacity_bytes() -> Option { + parse_alloc_cache_capacity(std::env::var("ARC_ALLOC_CACHE_MAX_MB").ok().as_deref()) +} + +/// Pure half of [`alloc_cache_capacity_bytes`], so the parse is testable. +/// +/// `None` (unset, or unparseable) leaves candle's own default in force rather +/// than silently picking a different one — a typo in an ops script must not +/// quietly hand back the unbounded allocator this change exists to remove. +pub(crate) fn parse_alloc_cache_capacity(raw: Option<&str>) -> Option { + let mb: usize = raw?.trim().parse().ok()?; + Some(if mb == 0 { + usize::MAX + } else { + mb.saturating_mul(1024 * 1024) + }) +} + +/// Emit the allocator's counters every `ARC_ALLOC_CACHE_STATS` decode steps. +/// +/// The counters are what this cache has to be judged on. A green log — the +/// server did not crash, the output looked fine — is not evidence that the +/// cache is working: an earlier arena in this codebase reported "accounting OK" +/// and bit-identical output across 52 steps while silently bypassing itself for +/// every buffer under 128 bytes (`KERNEL_RULES.md:977-984`). What distinguishes +/// a working bounded cache from that is arithmetic, and it is printed here: +/// +/// * `alloc/step` — real `cuMemAllocAsync` calls per decode step. Low means the +/// cache is absorbing the ~11 k allocations a step makes. +/// * `free/step` — real `cuMemFreeAsync` calls per decode step. **Non-zero is +/// the point.** The unbounded cache's was exactly zero, forever, which is why +/// it grew without bound. +/// * `held` — bytes retained right now, against the cap. +#[cfg(feature = "cuda")] +fn report_alloc_cache_step(cd: &candle_core::CudaDevice, seq_len: usize) { + use std::sync::atomic::{AtomicU64, Ordering}; + static EVERY: std::sync::OnceLock = std::sync::OnceLock::new(); + let every = *EVERY.get_or_init(|| { + std::env::var("ARC_ALLOC_CACHE_STATS") + .ok() + .and_then(|v| v.trim().parse::().ok()) + .unwrap_or(0) + }); + if every == 0 || seq_len != 1 { + return; + } + static STEP: AtomicU64 = AtomicU64::new(0); + static LAST_ALLOC: AtomicU64 = AtomicU64::new(0); + static LAST_FREE: AtomicU64 = AtomicU64::new(0); + static LAST_STEP: AtomicU64 = AtomicU64::new(0); + let step = STEP.fetch_add(1, Ordering::Relaxed) + 1; + if step % every != 0 { + return; + } + let s = cd.alloc_cache_stats(); + let d_step = step - LAST_STEP.swap(step, Ordering::Relaxed); + let d_alloc = s.misses - LAST_ALLOC.swap(s.misses, Ordering::Relaxed); + let d_free = s.frees() - LAST_FREE.swap(s.frees(), Ordering::Relaxed); + let n = d_step.max(1) as f64; + tracing::info!( + "[alloc-cache] step {step} alloc/step {:.1} free/step {:.1} \ + hit-rate {:.4} held {:.1} MiB / cap {} high-water {:.1} MiB sizes {}", + d_alloc as f64 / n, + d_free as f64 / n, + s.hits as f64 / (s.hits + s.misses).max(1) as f64, + s.cached_bytes as f64 / (1024.0 * 1024.0), + if s.capacity_bytes == usize::MAX { + "unbounded".to_string() + } else { + format!("{} MiB", s.capacity_bytes / (1024 * 1024)) + }, + s.high_water_bytes as f64 / (1024.0 * 1024.0), + s.size_classes, + ); +} + /// A loader for a "normal" (non-quantized) model. pub struct NormalLoader { inner: Box, @@ -1804,12 +1888,18 @@ impl Pipeline for NormalPipeline { cd.alloc_cache_enabled(), alloc_cache_killed(), ) { - AllocCacheAction::Enable => cd.set_alloc_cache_enabled(true), + AllocCacheAction::Enable => { + cd.set_alloc_cache_enabled(true); + if let Some(cap) = alloc_cache_capacity_bytes() { + cd.set_alloc_cache_capacity(cap); + } + } AllocCacheAction::DrainAndDisable => { cd.set_alloc_cache_enabled(false) } AllocCacheAction::Leave => {} } + report_alloc_cache_step(cd, seq_len); } if probe && seq_len == 1 { // RUN-161 step 2b. Set the graph-mode device position: @@ -2712,3 +2802,49 @@ mod alloc_cache_policy_tests { ); } } + +/// The retention cap's parse. +/// +/// The cap is the whole of the fix — candle's caching allocator had none, and +/// retained 6.04 MiB per decoded token forever — so the way it is configured has +/// to be unambiguous. In particular a typo must not fall back to "unbounded". +#[cfg(test)] +mod alloc_cache_capacity_tests { + use super::parse_alloc_cache_capacity; + + #[test] + fn unset_leaves_candles_default_in_force() { + assert_eq!(parse_alloc_cache_capacity(None), None); + } + + /// A typo is not a licence to run unbounded. `None` means "don't touch it", + /// and candle's default is bounded, so a bad value degrades to bounded. + #[test] + fn an_unparseable_value_does_not_become_unbounded() { + for bad in ["", " ", "lots", "1GiB", "-1", "1.5"] { + assert_eq!(parse_alloc_cache_capacity(Some(bad)), None, "{bad:?}"); + } + } + + /// `0` is the documented escape hatch back to the old unbounded allocator, + /// kept reachable only so its leak can be re-measured rather than argued + /// about. + #[test] + fn zero_is_the_explicit_unbounded_opt_in() { + assert_eq!(parse_alloc_cache_capacity(Some("0")), Some(usize::MAX)); + } + + #[test] + fn megabytes_convert_and_do_not_overflow() { + assert_eq!(parse_alloc_cache_capacity(Some("1")), Some(1024 * 1024)); + assert_eq!( + parse_alloc_cache_capacity(Some(" 512 ")), + Some(512 * 1024 * 1024) + ); + assert_eq!( + parse_alloc_cache_capacity(Some(&usize::MAX.to_string())), + Some(usize::MAX), + "saturates instead of wrapping to a tiny cap" + ); + } +} From f9d496acf929e015bebc43b2b292e58ceebd0c82 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Wed, 19 Aug 2026 23:54:30 +0100 Subject: [PATCH 14/22] fix(build): report_alloc_cache_step takes a reference Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01SpVNMpb13HkUXqSqbN1o9H --- mistralrs-core/src/pipeline/normal.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mistralrs-core/src/pipeline/normal.rs b/mistralrs-core/src/pipeline/normal.rs index a69835923..4675ab926 100644 --- a/mistralrs-core/src/pipeline/normal.rs +++ b/mistralrs-core/src/pipeline/normal.rs @@ -1899,7 +1899,7 @@ impl Pipeline for NormalPipeline { } AllocCacheAction::Leave => {} } - report_alloc_cache_step(cd, seq_len); + report_alloc_cache_step(&cd, seq_len); } if probe && seq_len == 1 { // RUN-161 step 2b. Set the graph-mode device position: From fa1d2ef28b1def98d3767b7302018e862ddc66e4 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Thu, 20 Aug 2026 00:22:50 +0100 Subject: [PATCH 15/22] =?UTF-8?q?chore(deps):=20candle=20859c49c8=20?= =?UTF-8?q?=E2=80=94=20allocator=20test=20fixes?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Test-only change on the candle side; `git diff b2a4dd80..859c49c8` touches nothing outside `#[cfg(test)] mod alloc_cache_tests`. The H200 numbers in the branch description were measured at b2a4dd80 and stand. Correction to that commit message: the allocator has **ten** tests, not eleven. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01SpVNMpb13HkUXqSqbN1o9H --- Cargo.toml | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 225d8ecbe..0f7f7935f 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -49,11 +49,11 @@ rust-version = "1.88" # pointer has broken CI on this repo before — a path/branch dependency that # advances under a merged PR turns a green build into a red one with no commit # here to blame. -candle-core = { git = "https://github.com/aeonmindai/candle.git", rev = "b2a4dd80b82726730c3c4c90421fa3a0938be984" } -candle-nn = { git = "https://github.com/aeonmindai/candle.git", rev = "b2a4dd80b82726730c3c4c90421fa3a0938be984" } -candle-flash-attn-v3 = { git = "https://github.com/aeonmindai/candle.git", rev = "b2a4dd80b82726730c3c4c90421fa3a0938be984" } -candle-flash-attn = { git = "https://github.com/aeonmindai/candle.git", rev = "b2a4dd80b82726730c3c4c90421fa3a0938be984" } -candle-metal-kernels = { git = "https://github.com/aeonmindai/candle.git", rev = "b2a4dd80b82726730c3c4c90421fa3a0938be984" } +candle-core = { git = "https://github.com/aeonmindai/candle.git", rev = "859c49c8c4e24177378a816f7b3f3d284987c62f" } +candle-nn = { git = "https://github.com/aeonmindai/candle.git", rev = "859c49c8c4e24177378a816f7b3f3d284987c62f" } +candle-flash-attn-v3 = { git = "https://github.com/aeonmindai/candle.git", rev = "859c49c8c4e24177378a816f7b3f3d284987c62f" } +candle-flash-attn = { git = "https://github.com/aeonmindai/candle.git", rev = "859c49c8c4e24177378a816f7b3f3d284987c62f" } +candle-metal-kernels = { git = "https://github.com/aeonmindai/candle.git", rev = "859c49c8c4e24177378a816f7b3f3d284987c62f" } axum = "0.8.8" anyhow = "1.0.100" From 65dc0fb4a36d4615a477afad5cd9ef2f25278c3b Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Thu, 20 Aug 2026 00:50:50 +0100 Subject: [PATCH 16/22] =?UTF-8?q?chore(deps):=20candle=2089ab14ef=20?= =?UTF-8?q?=E2=80=94=20no=20eviction=20inside=20the=20capture=20window?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Closes a free path this change introduced: leaving capture mode re-filed parked buffers through the evicting put, which could hand a private-pool pointer back to the driver at a moment arc-cuda-graph does not control. Decode is unaffected — `set_capture_mode` is only called during a capture. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01SpVNMpb13HkUXqSqbN1o9H --- Cargo.toml | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index 0f7f7935f..b8852835d 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -49,11 +49,11 @@ rust-version = "1.88" # pointer has broken CI on this repo before — a path/branch dependency that # advances under a merged PR turns a green build into a red one with no commit # here to blame. -candle-core = { git = "https://github.com/aeonmindai/candle.git", rev = "859c49c8c4e24177378a816f7b3f3d284987c62f" } -candle-nn = { git = "https://github.com/aeonmindai/candle.git", rev = "859c49c8c4e24177378a816f7b3f3d284987c62f" } -candle-flash-attn-v3 = { git = "https://github.com/aeonmindai/candle.git", rev = "859c49c8c4e24177378a816f7b3f3d284987c62f" } -candle-flash-attn = { git = "https://github.com/aeonmindai/candle.git", rev = "859c49c8c4e24177378a816f7b3f3d284987c62f" } -candle-metal-kernels = { git = "https://github.com/aeonmindai/candle.git", rev = "859c49c8c4e24177378a816f7b3f3d284987c62f" } +candle-core = { git = "https://github.com/aeonmindai/candle.git", rev = "89ab14ef1216331a782539f533b99f6816708604" } +candle-nn = { git = "https://github.com/aeonmindai/candle.git", rev = "89ab14ef1216331a782539f533b99f6816708604" } +candle-flash-attn-v3 = { git = "https://github.com/aeonmindai/candle.git", rev = "89ab14ef1216331a782539f533b99f6816708604" } +candle-flash-attn = { git = "https://github.com/aeonmindai/candle.git", rev = "89ab14ef1216331a782539f533b99f6816708604" } +candle-metal-kernels = { git = "https://github.com/aeonmindai/candle.git", rev = "89ab14ef1216331a782539f533b99f6816708604" } axum = "0.8.8" anyhow = "1.0.100" From e4eb59dfba78953ee8d85aa1dc3da69801e6452c Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Thu, 20 Aug 2026 11:56:53 +0100 Subject: [PATCH 17/22] fix(ci): spell 'unparsable' the way the typos gate expects MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Two sites in pipeline/normal.rs — a doc comment and a test name. No behaviour change; the Typos job was the only red check on this PR. Co-Authored-By: Claude Opus 5 (1M context) --- mistralrs-core/src/pipeline/normal.rs | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/mistralrs-core/src/pipeline/normal.rs b/mistralrs-core/src/pipeline/normal.rs index 4675ab926..0bcc24e8a 100644 --- a/mistralrs-core/src/pipeline/normal.rs +++ b/mistralrs-core/src/pipeline/normal.rs @@ -197,7 +197,7 @@ pub(crate) fn alloc_cache_capacity_bytes() -> Option { /// Pure half of [`alloc_cache_capacity_bytes`], so the parse is testable. /// -/// `None` (unset, or unparseable) leaves candle's own default in force rather +/// `None` (unset, or unparsable) leaves candle's own default in force rather /// than silently picking a different one — a typo in an ops script must not /// quietly hand back the unbounded allocator this change exists to remove. pub(crate) fn parse_alloc_cache_capacity(raw: Option<&str>) -> Option { @@ -2820,7 +2820,7 @@ mod alloc_cache_capacity_tests { /// A typo is not a licence to run unbounded. `None` means "don't touch it", /// and candle's default is bounded, so a bad value degrades to bounded. #[test] - fn an_unparseable_value_does_not_become_unbounded() { + fn an_unparsable_value_does_not_become_unbounded() { for bad in ["", " ", "lots", "1GiB", "-1", "1.5"] { assert_eq!(parse_alloc_cache_capacity(Some(bad)), None, "{bad:?}"); } From d71fec0278b231770097655cd79b5b96333d7e5f Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick <48842933+heydryft@users.noreply.github.com> Date: Thu, 20 Aug 2026 18:39:27 +0100 Subject: [PATCH 18/22] fix(deps): regenerate Cargo.lock for the candle 89ab14ef pin `chore(deps): candle 89ab14ef` bumped all five candle crates in `Cargo.toml` but left `Cargo.lock` pinning 88d86a2. That is the stale-lock trap: a `--locked` CI build resolves the LOCK, so the lane would have compiled the OLD, UNBOUNDED allocator while reporting green on a PR whose entire subject is bounding it. `cargo metadata --locked` now succeeds; it failed before this commit. Regenerated with `cargo update -p candle-{core,nn,flash-attn,flash-attn-v3,metal-kernels}`. The only non-candle churn is a `windows-core` 0.61.2/0.62.2 dedup that fell out of the re-resolution; nothing arc builds on Linux or macOS reads it. --- Cargo.lock | 57 +++++++++++++++++++++--------------------------------- 1 file changed, 22 insertions(+), 35 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 1e5fa1e5f..5a68f4570 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -132,7 +132,7 @@ version = "1.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" dependencies = [ - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -143,7 +143,7 @@ checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" dependencies = [ "anstyle", "once_cell_polyfill", - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -692,7 +692,7 @@ checksum = "ade8366b8bd5ba243f0a58f036cc0ca8a2f069cff1a2351ef1cac6b083e16fc0" [[package]] name = "candle-core" version = "0.9.2" -source = "git+https://github.com/aeonmindai/candle.git?rev=88d86a2019ee552923b217284e89847d2785dcdf#88d86a2019ee552923b217284e89847d2785dcdf" +source = "git+https://github.com/aeonmindai/candle.git?rev=89ab14ef1216331a782539f533b99f6816708604#89ab14ef1216331a782539f533b99f6816708604" dependencies = [ "accelerate-src", "byteorder", @@ -723,7 +723,7 @@ dependencies = [ [[package]] name = "candle-flash-attn" version = "0.9.2" -source = "git+https://github.com/aeonmindai/candle.git?rev=88d86a2019ee552923b217284e89847d2785dcdf#88d86a2019ee552923b217284e89847d2785dcdf" +source = "git+https://github.com/aeonmindai/candle.git?rev=89ab14ef1216331a782539f533b99f6816708604#89ab14ef1216331a782539f533b99f6816708604" dependencies = [ "anyhow", "candle-core", @@ -734,7 +734,7 @@ dependencies = [ [[package]] name = "candle-flash-attn-v3" version = "0.9.2" -source = "git+https://github.com/aeonmindai/candle.git?rev=88d86a2019ee552923b217284e89847d2785dcdf#88d86a2019ee552923b217284e89847d2785dcdf" +source = "git+https://github.com/aeonmindai/candle.git?rev=89ab14ef1216331a782539f533b99f6816708604#89ab14ef1216331a782539f533b99f6816708604" dependencies = [ "anyhow", "candle-core", @@ -747,7 +747,7 @@ dependencies = [ [[package]] name = "candle-kernels" version = "0.9.2" -source = "git+https://github.com/aeonmindai/candle.git?rev=88d86a2019ee552923b217284e89847d2785dcdf#88d86a2019ee552923b217284e89847d2785dcdf" +source = "git+https://github.com/aeonmindai/candle.git?rev=89ab14ef1216331a782539f533b99f6816708604#89ab14ef1216331a782539f533b99f6816708604" dependencies = [ "cudaforge", ] @@ -755,7 +755,7 @@ dependencies = [ [[package]] name = "candle-metal-kernels" version = "0.9.2" -source = "git+https://github.com/aeonmindai/candle.git?rev=88d86a2019ee552923b217284e89847d2785dcdf#88d86a2019ee552923b217284e89847d2785dcdf" +source = "git+https://github.com/aeonmindai/candle.git?rev=89ab14ef1216331a782539f533b99f6816708604#89ab14ef1216331a782539f533b99f6816708604" dependencies = [ "half", "objc2", @@ -769,7 +769,7 @@ dependencies = [ [[package]] name = "candle-nn" version = "0.9.2" -source = "git+https://github.com/aeonmindai/candle.git?rev=88d86a2019ee552923b217284e89847d2785dcdf#88d86a2019ee552923b217284e89847d2785dcdf" +source = "git+https://github.com/aeonmindai/candle.git?rev=89ab14ef1216331a782539f533b99f6816708604#89ab14ef1216331a782539f533b99f6816708604" dependencies = [ "accelerate-src", "candle-core", @@ -788,7 +788,7 @@ dependencies = [ [[package]] name = "candle-ug" version = "0.9.2" -source = "git+https://github.com/aeonmindai/candle.git?rev=88d86a2019ee552923b217284e89847d2785dcdf#88d86a2019ee552923b217284e89847d2785dcdf" +source = "git+https://github.com/aeonmindai/candle.git?rev=89ab14ef1216331a782539f533b99f6816708604#89ab14ef1216331a782539f533b99f6816708604" dependencies = [ "ug", "ug-cuda", @@ -1633,7 +1633,7 @@ dependencies = [ "libc", "option-ext", "redox_users 0.5.2", - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -1813,7 +1813,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -2759,7 +2759,7 @@ dependencies = [ "js-sys", "log", "wasm-bindgen", - "windows-core 0.62.2", + "windows-core", ] [[package]] @@ -4091,7 +4091,7 @@ version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -5483,7 +5483,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys 0.12.1", - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -5542,7 +5542,7 @@ dependencies = [ "security-framework", "security-framework-sys", "webpki-root-certs", - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -6068,7 +6068,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3a766e1110788c36f4fa1c2b71b387a7815aa65f88ce0229841826633d93723e" dependencies = [ "libc", - "windows-sys 0.61.2", + "windows-sys 0.60.2", ] [[package]] @@ -6515,7 +6515,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "230a1b821ccbd75b185820a1f1ff7b14d21da1e442e22c0863ea5f08771a8874" dependencies = [ "rustix 1.1.4", - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -7626,7 +7626,7 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" dependencies = [ - "windows-sys 0.61.2", + "windows-sys 0.59.0", ] [[package]] @@ -7642,7 +7642,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9babd3a767a4c1aef6900409f85f5d53ce2544ccdfaa86dad48c91782c6d6893" dependencies = [ "windows-collections", - "windows-core 0.61.2", + "windows-core", "windows-future", "windows-link 0.1.3", "windows-numerics", @@ -7654,7 +7654,7 @@ version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3beeceb5e5cfd9eb1d76b381630e82c4241ccd0d27f1a39ed41b2760b255c5e8" dependencies = [ - "windows-core 0.61.2", + "windows-core", ] [[package]] @@ -7670,26 +7670,13 @@ dependencies = [ "windows-strings 0.4.2", ] -[[package]] -name = "windows-core" -version = "0.62.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b8e83a14d34d0623b51dce9581199302a221863196a1dde71a7663a4c2be9deb" -dependencies = [ - "windows-implement", - "windows-interface", - "windows-link 0.2.1", - "windows-result 0.4.1", - "windows-strings 0.5.1", -] - [[package]] name = "windows-future" version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc6a41e98427b19fe4b73c550f060b59fa592d7d686537eebf9385621bfbad8e" dependencies = [ - "windows-core 0.61.2", + "windows-core", "windows-link 0.1.3", "windows-threading", ] @@ -7734,7 +7721,7 @@ version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9150af68066c4c5c07ddc0ce30421554771e528bde427614c61038bc2c92c2b1" dependencies = [ - "windows-core 0.61.2", + "windows-core", "windows-link 0.1.3", ] From 9f110905b070e97977f91eef337b8a5fdf826730 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick Date: Wed, 19 Aug 2026 22:01:53 +0100 Subject: [PATCH 19/22] =?UTF-8?q?perf(ArcMoE):=20fuse=20the=20V4=20router?= =?UTF-8?q?=20region=20=E2=80=94=2018+9=20launches=20become=202,=20and=20f?= =?UTF-8?q?ix=20the=20sinkhorn=20local-memory=20spill?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit At b=1 the V4 decode step is launch-bound: ~7,924 cuLaunchKernel calls, of which the MoE router region is 3,259 (41.8%) across 86 contiguous spans — exactly two per layer (mhc_attn_pre, mhc_ffn_pre; 43 layers). Mean kernel in that region is 2.79 us and 89% are under 5 us, on [1,24] / [1,4,4] / [1,256] tensors. That is launch overhead wearing a kernel costume. Three changes, all bit-identical to what they replace: 1. cuda/hc_fused.cu `hc_pre_fused_f32` collapses hc_pre's 18-launch middle — a hand-decomposed RMS statistic in SEVEN launches (sqr, fast_sum, affine, affine, recip, sqrt, bmul) plus ELEVEN for the pre/post/comb scoring — into one kernel. `mixes` is never materialised. 2. cuda/hc_fused.cu `sqrt_softplus_f32` collapses MoeGate's sqrt(softplus(x)) from nine launches (zeros_like, bmaximum, uabs, uneg, uexp, affine, ulog, badd, usqrt) into one. 3. cuda/sinkhorn.cu is templated on `hc`. It took `hc` at runtime, so nvcc could not keep r[16]/buf[16]/col[16] in registers and demoted all three to local memory — `ptxas -v` reported 192 bytes stack frame. With hc=4 the kernel runs ONE block of FOUR threads, so nothing hides that latency: measured 30.0 us/call on a [1,4,4] tensor, 2.549 ms/step over 86 calls, 28% of the router's GPU time. Templating makes every index compile-time; ptxas now reports 0 bytes stack frame, 29 registers. Arithmetic unchanged. Bit-identity is by construction, not by tolerance: this region decides WHICH EXPERTS RUN, so a reassociated sum or a contracted FMA can change the emitted token. hc_fused.cu joins sinkhorn.cu in build.rs's dedicated no-fast-math, --fmad=false builder and carries the same #error guard, and it transcribes candle's ops exactly — including that `recipg` is `1.0 / a` with a *double* literal, that `mean_keepdim` is fast_sum followed by a separate affine, and that candle's FastReduce block_dim is min(1024, len).next_power_of_two(), which fixes the reduction tree's shape and therefore the f32 rounding. cuda/hc_fused.rs carries scalar replicas of both sides asserted bit-identical without a GPU, plus two anti-vacuity guards. The softplus one earned its place: its first version sampled [-30, -12.5, -1, 0.5, 3.25, 17] and passed against the naive log(1+exp(x)) form, because the two agree bit for bit until exp overflows near 88.7. It is now a three-way mutation probe. ARC_HC_FUSED=0 restores the eager chains so both can be A/B'd from one binary. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01SpVNMpb13HkUXqSqbN1o9H --- mistralrs-core/build.rs | 27 +- mistralrs-core/src/cuda/ffi.rs | 28 ++ mistralrs-core/src/cuda/hc_fused.cu | 291 ++++++++++++ mistralrs-core/src/cuda/hc_fused.rs | 588 +++++++++++++++++++++++++ mistralrs-core/src/cuda/mod.rs | 1 + mistralrs-core/src/cuda/sinkhorn.cu | 203 ++++++--- mistralrs-core/src/models/deepseek4.rs | 20 +- mistralrs-core/src/models/dsv4_mhc.rs | 136 ++++-- 8 files changed, 1180 insertions(+), 114 deletions(-) create mode 100644 mistralrs-core/src/cuda/hc_fused.cu create mode 100644 mistralrs-core/src/cuda/hc_fused.rs diff --git a/mistralrs-core/build.rs b/mistralrs-core/build.rs index ef506a3db..bfe9f4226 100644 --- a/mistralrs-core/build.rs +++ b/mistralrs-core/build.rs @@ -10,15 +10,15 @@ fn main() { println!("cargo:rerun-if-changed=build.rs"); let build_dir = PathBuf::from(std::env::var("OUT_DIR").unwrap()); - // sinkhorn.cu is EXCLUDED from this fast-math builder and compiled - // separately below: it must be bit-identical to candle-kernels (which - // build with plain -O3, no fast math), and --use_fast_math rewrites - // expf -> __expf and IEEE division -> approximate reciprocals. The - // kernel source carries an `#error` guard against being re-globbed - // under fast math. See mistralrs-core/src/cuda/sinkhorn.cu. + // sinkhorn.cu and hc_fused.cu are EXCLUDED from this fast-math builder + // and compiled separately below: they must be bit-identical to + // candle-kernels (which build with plain -O3, no fast math), and + // --use_fast_math rewrites expf -> __expf and IEEE division -> + // approximate reciprocals. Both sources carry an `#error` guard against + // being re-globbed under fast math. See mistralrs-core/src/cuda/. let mut builder = cudaforge::KernelBuilder::new() .source_glob("src/cuda/*.cu") - .exclude(&["sinkhorn.cu"]) + .exclude(&["sinkhorn.cu", "hc_fused.cu"]) .out_dir(&build_dir) .arg("-std=c++17") .arg("-O3") @@ -63,14 +63,15 @@ fn main() { println!("cargo:rustc-link-search={}", build_dir.display()); println!("cargo:rustc-link-lib=mistralrscuda"); - // Dedicated IEEE (no fast math) builder for sinkhorn.cu — bit-identity - // with candle-kernels requires accurate expf + div.rn.f32; --fmad=false - // additionally forbids FMA contraction so rounding matches candle's - // unfused op chain exactly. Own subdirectory so its build cache never - // mixes with the fast-math builder's. + // Dedicated IEEE (no fast math) builder for the bit-identity-critical + // kernels — bit-identity with candle-kernels requires accurate + // expf/logf + div.rn.f32; --fmad=false additionally forbids FMA + // contraction so rounding matches candle's unfused op chain exactly. + // Own subdirectory so its build cache never mixes with the fast-math + // builder's. let sinkhorn_dir = build_dir.join("sinkhorn_ieee"); let mut sinkhorn_builder = cudaforge::KernelBuilder::new() - .source_files(vec!["src/cuda/sinkhorn.cu"]) + .source_files(vec!["src/cuda/sinkhorn.cu", "src/cuda/hc_fused.cu"]) .out_dir(&sinkhorn_dir) .arg("-std=c++17") .arg("-O3") diff --git a/mistralrs-core/src/cuda/ffi.rs b/mistralrs-core/src/cuda/ffi.rs index 81757e1b7..a2ca78e45 100644 --- a/mistralrs-core/src/cuda/ffi.rs +++ b/mistralrs-core/src/cuda/ffi.rs @@ -371,4 +371,32 @@ extern "C" { eps: f32, stream: i64, ); + // cuda/hc_fused.cu — fused V4 router-region kernels. Like sinkhorn.cu these + // are compiled by build.rs's dedicated no-fast-math builder; see the + // bit-identity contract at the top of that file. + #[allow(clippy::too_many_arguments)] + pub(crate) fn hc_pre_fused_f32( + x_flat: *const c_void, + mixes_raw: *const c_void, + hc_scale: *const c_void, + hc_base: *const c_void, + pre: *mut c_void, + post: *mut c_void, + comb_pre: *mut c_void, + n: i32, + d: i32, + m: i32, + hc: i32, + block_dim: i32, + inv_d: f32, + rms_eps: f32, + hc_eps: f32, + stream: i64, + ); + pub(crate) fn sqrt_softplus_f32( + inp: *const c_void, + out: *mut c_void, + numel: i64, + stream: i64, + ); } diff --git a/mistralrs-core/src/cuda/hc_fused.cu b/mistralrs-core/src/cuda/hc_fused.cu new file mode 100644 index 000000000..bdfc711ce --- /dev/null +++ b/mistralrs-core/src/cuda/hc_fused.cu @@ -0,0 +1,291 @@ +// Parent system: ArcInfer / ArcMoE +// +// Fused router-region kernels for DeepSeek-V4 (mHC pre-step + MoE gate scoring). +// +// --------------------------------------------------------------------------- +// WHY +// --------------------------------------------------------------------------- +// At b=1 the V4 decode step is launch-bound, not bandwidth-bound: the measured +// step issues ~7.9k `cuLaunchKernel` calls, of which the router region alone is +// 3,259 (41.8%) spread over 86 contiguous spans -- exactly two per layer +// (`mhc_attn_pre` and `mhc_ffn_pre`, 43 layers). Those spans operate on +// `[1, 24]` / `[1, 4, 4]` / `[1, 256]` tensors: mean kernel duration 2.79 us, +// 89% of them under 5 us. They are pure launch overhead wearing a kernel +// costume. +// +// Two expressions dominate the count and are collapsed here: +// +// 1. `hc_pre` (dsv4_mhc.rs) spells a hand-decomposed RMS statistic in SEVEN +// launches -- `sqr -> fast_sum -> affine(1/D) -> affine(+eps) -> urecip -> +// usqrt -> bmul` -- and then a further ELEVEN for the pre/post/comb +// sigmoid-and-bias scoring. Eighteen launches whose entire data footprint +// after the reduction is 24 floats. +// +// 2. `MoeGate::forward` (deepseek4.rs) spells `sqrt(softplus(x))` as NINE +// launches: `zeros_like -> bmaximum -> uabs -> uneg -> uexp -> affine(+1) +// -> ulog -> badd -> usqrt`. +// +// --------------------------------------------------------------------------- +// BIT-IDENTITY CONTRACT +// --------------------------------------------------------------------------- +// The mHC pre-step feeds `y` into attention and the gate scores decide WHICH +// EXPERTS RUN. A reassociated sum or a contracted FMA can flip an expert choice +// and therefore the generated token, so these kernels are bit-identical to the +// candle op chains they replace, BY CONSTRUCTION, not by tolerance. The same +// three rules that `sinkhorn.cu` documents apply verbatim: +// +// 1. REDUCTION ORDER. candle's `sum_keepdim` on CUDA runs candle-kernels +// `fast_sum` (reduce.cu) with `block_dim = min(1024, el_to_sum_per_block) +// .next_power_of_two()` (cuda_backend/mod.rs `FastReduce`). Each thread +// accumulates a *strided* slice sequentially into a zero-initialized +// accumulator (`shr[tid] = 0; shr[tid] += src[idx]; idx += blockDim.x`), +// then a pairwise tree `shr[t] += shr[t + s]` for s = block/2 .. 1. For +// D = hc_mult * hidden = 16384 that is 1024 threads x 16 elements each, and +// the tree order is NOT the sequential order. `hc_pre_fused_f32_kernel` +// replays it exactly, including the `shr[tid] = 0` initialisation (which +// canonicalises -0.0f to +0.0f) and the identity padding when a thread's +// strided slice is empty. +// +// 2. UNFUSED, IEEE ROUND-TO-NEAREST ARITHMETIC. candle-kernels build with +// plain -O3 and NO --use_fast_math, so their division is `div.rn.f32`, +// their `expf`/`logf` are the accurate libdevice `__nv_expf`/`__nv_logf`, +// and denormals are not flushed. mistralrs-core's build.rs compiles the +// rest of src/cuda/*.cu WITH --use_fast_math, which would silently rewrite +// all three. THIS FILE IS THEREFORE COMPILED BY THE SAME DEDICATED +// no-fast-math / --fmad=false builder that sinkhorn.cu uses, and carries +// the same #error guard against a future re-glob. All arithmetic uses +// __fadd_rn / __fmul_rn / __fsub_rn, which nvcc documents as never merged +// into an FMA, so the fusion cannot silently change rounding. +// +// 3. EXACT OP TRANSCRIPTION, including the ones that look like no-ops: +// - candle's `recipg(float a)` is `return 1.0 / a;` -- the literal is a +// *double*, so the op is `(float)(1.0 / (double)a)`, not `__frcp_rn(a)`. +// Transcribed literally in `candle_recip`. +// - `Tensor * f64` / `Tensor + f64` lower to `affine(mul, add)` whose +// kernel is `x * mul + add`, compiled by candle with the default +// -fmad=true, i.e. `fmaf(x, mul, add)`. Transcribed as explicit `fmaf` +// (an explicit fmaf is unaffected by our --fmad=false, which only +// governs *contraction* of separate mul+add). +// - `mean_keepdim` is `sum_impl(..)? * (1f64 / len as f64)`, i.e. a +// `fast_sum` followed by a SEPARATE affine -- the division is not folded +// into the reduction. +// - `candle_nn::ops::sigmoid` dispatches to candle-kernels `usigmoid_f32` +// = `sigmoid_fwd(x)` = `recipg(1.0f + expf(-x))`. +// +// Scalar Rust replicas of both sides (the candle op chain and this kernel) are +// asserted bit-identical over randomized inputs in `cuda/hc_fused.rs` +// `mod tests`, so a future edit to either side trips CPU CI without a GPU. The +// final proof is the on-GPU A/B: `ARC_HC_AB=1` recomputes the candle chain +// alongside the fused kernel and reports any bitwise mismatch. + +#include +#include + +#if defined(__USE_FAST_MATH__) +#error "hc_fused.cu must be compiled WITHOUT --use_fast_math: fast math rewrites expf/logf to the hardware approximations and IEEE division to approximate reciprocals, breaking bit-identity with candle-kernels (which build with plain -O3). See the dedicated no-fast-math builder in mistralrs-core/build.rs." +#endif + +namespace { + +// candle-kernels cuda_utils.cuh: +// __device__ __forceinline__ float recipg(float a) { return 1.0 / a; } +// The literal `1.0` is a double, so `a` is promoted, the division is done in +// double, and the result is narrowed on return. Transcribed literally rather +// than "simplified" to __frcp_rn: the two happen to agree for all finite +// inputs, but the point of this file is to not rely on such arguments. +__device__ __forceinline__ float candle_recip(float a) { + return (float)(1.0 / (double)a); +} + +// candle-kernels unary.cu: +// __device__ __forceinline__ T sigmoid_fwd(T x) { +// return recipg(static_cast(1) + expg(-x)); +// } +// with expg(float) == expf (libdevice __nv_expf under no-fast-math). +__device__ __forceinline__ float candle_sigmoid(float x) { + return candle_recip(__fadd_rn(1.0f, expf(-x))); +} + +} // namespace + +extern "C" { + +// --------------------------------------------------------------------------- +// hc_pre: the 18-launch RMS-scale + pre/post/comb scoring block, in one launch. +// --------------------------------------------------------------------------- +// +// Replaces, verbatim (dsv4_mhc.rs `hc_pre`): +// +// let sq_mean = x_flat.sqr()?.mean_keepdim(D::Minus1)?; // 3 +// let rsqrt = (sq_mean + rms_norm_eps)?.recip()?.sqrt()?; // 3 +// let mixes = mixes_raw.broadcast_mul(&rsqrt)?; // 1 +// let pre = sigmoid(pre_block *s_pre +b_pre)? + hc_eps; // 4 +// let post = sigmoid(post_block*s_post +b_post)?.affine(2,0); // 4 +// let comb_pre= comb_block*s_comb + b_comb; // 3 +// +// The trailing 3 for `comb` include the `.reshape((n, hc, hc))` of a narrowed +// (non-contiguous) view, which candle services with a real copy kernel; here +// the reshape is implicit in the output indexing and costs nothing. +// +// `mixes` itself is never materialised: it is consumed only by the three +// narrows, so the fused kernel keeps it in registers. +// +// Launch contract (enforced host-side in hc_fused.rs): +// grid = n blocks (one per row, matching candle's `dst_el` grid) +// block = min(1024, d).next_power_of_two() -- candle's FastReduce block_dim +// shmem = blockDim.x * sizeof(float) +// m <= blockDim.x, and x_flat / mixes_raw are contiguous F32. +__global__ void hc_pre_fused_f32_kernel( + const float *__restrict__ x_flat, // [n, d] F32 contiguous + const float *__restrict__ mixes_raw, // [n, m] F32 contiguous + const float *__restrict__ hc_scale, // [3] F32 + const float *__restrict__ hc_base, // [m] F32 + float *__restrict__ pre, // [n, hc] out + float *__restrict__ post, // [n, hc] out + float *__restrict__ comb_pre, // [n, hc, hc] out + int d, + int m, + int hc, + float inv_d, // (float)(1.0 / d) -- candle's mean scale, T::from_f64 + float rms_eps, // (float) rms_norm_eps + float hc_eps // (float) hc_eps +) { + extern __shared__ float shr[]; + const int n = blockIdx.x; + const int tid = threadIdx.x; + const int bd = blockDim.x; + + // ---- candle `sqr()` + `fast_sum` replay ------------------------------- + // candle runs these as two kernels; the intermediate `x*x` is rounded to + // f32 before the accumulation, which __fmul_rn reproduces. The strided + // walk and the zero-initialised accumulator are fast_sum's, verbatim. + const float *xrow = x_flat + (size_t)n * (size_t)d; + float acc = 0.0f; + for (int idx = tid; idx < d; idx += bd) { + const float v = xrow[idx]; + acc = __fadd_rn(acc, __fmul_rn(v, v)); + } + shr[tid] = acc; + // fast_sum's pairwise tree. The __syncthreads() sits at the TOP of the + // body there too, so the write above is visible before the first read. + for (int s = bd >> 1; s > 0; s >>= 1) { + __syncthreads(); + if (tid < s) { + shr[tid] = __fadd_rn(shr[tid], shr[tid + s]); + } + } + __syncthreads(); + + // ---- mean_keepdim tail, then `+eps -> recip -> sqrt` ------------------ + // Every thread recomputes this from shr[0]; it is deterministic scalar + // arithmetic on one value, so all threads agree bit for bit. + const float sum = shr[0]; + const float mean = fmaf(sum, inv_d, 0.0f); // mean_keepdim: affine(1/d, 0) + const float meps = fmaf(mean, 1.0f, rms_eps); // (sq_mean + rms_norm_eps) + const float rsqrt = sqrtf(candle_recip(meps)); + + // ---- broadcast_mul into `mixes`, then the three scoring blocks -------- + if (tid < m) { + const float mx = __fmul_rn(mixes_raw[(size_t)n * (size_t)m + tid], rsqrt); + if (tid < hc) { + // pre = sigmoid(pre_block * s_pre + b_pre) + hc_eps + const float t = __fadd_rn(__fmul_rn(mx, hc_scale[0]), hc_base[tid]); + pre[(size_t)n * (size_t)hc + tid] = fmaf(candle_sigmoid(t), 1.0f, hc_eps); + } else if (tid < 2 * hc) { + // post = 2 * sigmoid(post_block * s_post + b_post) + const int j = tid - hc; + const float t = __fadd_rn(__fmul_rn(mx, hc_scale[1]), hc_base[hc + j]); + post[(size_t)n * (size_t)hc + j] = fmaf(candle_sigmoid(t), 2.0f, 0.0f); + } else { + // comb_pre = comb_block * s_comb + b_comb (reshape is free here) + const int j = tid - 2 * hc; + const float t = __fadd_rn(__fmul_rn(mx, hc_scale[2]), hc_base[2 * hc + j]); + comb_pre[(size_t)n * (size_t)hc * (size_t)hc + j] = t; + } + } +} + +void hc_pre_fused_f32( + const void *x_flat, + const void *mixes_raw, + const void *hc_scale, + const void *hc_base, + void *pre, + void *post, + void *comb_pre, + int n, + int d, + int m, + int hc, + int block_dim, + float inv_d, + float rms_eps, + float hc_eps, + long long stream +) { + dim3 grid(n, 1, 1); + dim3 block(block_dim, 1, 1); + size_t shmem = (size_t)block_dim * sizeof(float); + hc_pre_fused_f32_kernel<<>>( + (const float *)x_flat, + (const float *)mixes_raw, + (const float *)hc_scale, + (const float *)hc_base, + (float *)pre, + (float *)post, + (float *)comb_pre, + d, + m, + hc, + inv_d, + rms_eps, + hc_eps); +} + +// --------------------------------------------------------------------------- +// sqrt(softplus(x)): the V4 gate scoring function, in one launch instead of 9. +// --------------------------------------------------------------------------- +// +// Replaces, verbatim (deepseek4.rs `MoeGate::forward`, ScoringFunc::SqrtSoftplus): +// +// let max0 = logits.maximum(&logits.zeros_like()?)?; // zeros + bmaximum +// let abs = logits.abs()?; // uabs +// let softplus = (max0 + ((abs.neg()?.exp()? + 1.0)?.log()?))?; +// softplus.sqrt()? +// +// candle op-for-op: bmaximum is `maxg(x, y)` == `fmaxf`; uabs is `fabsf`; +// uneg is `-x`; uexp is `expf`; `+ 1.0` is `affine(1.0, 1.0)` == `fmaf(x,1,1)`; +// ulog is `logf`; the outer `+` is `badd` == `x + y`; usqrt is `sqrtf`. +__global__ void sqrt_softplus_f32_kernel( + const float *__restrict__ inp, + float *__restrict__ out, + long long numel +) { + const long long stride = (long long)blockDim.x * (long long)gridDim.x; + for (long long i = (long long)blockIdx.x * (long long)blockDim.x + threadIdx.x; + i < numel; + i += stride) { + const float l = inp[i]; + const float mx = fmaxf(l, 0.0f); // maximum(logits, zeros_like) + const float a = fabsf(l); // abs() + const float e = expf(-a); // neg() then exp() + const float p1 = fmaf(e, 1.0f, 1.0f); // + 1.0 (affine) + const float lg = logf(p1); // log() + out[i] = sqrtf(__fadd_rn(mx, lg)); // max0 + ... then sqrt() + } +} + +void sqrt_softplus_f32(const void *inp, void *out, long long numel, long long stream) { + const int block = 256; + long long blocks = (numel + block - 1) / block; + if (blocks < 1) { + blocks = 1; + } + if (blocks > 65535) { + blocks = 65535; + } + sqrt_softplus_f32_kernel<<<(unsigned)blocks, block, 0, (cudaStream_t)stream>>>( + (const float *)inp, (float *)out, numel); +} + +} // extern "C" diff --git a/mistralrs-core/src/cuda/hc_fused.rs b/mistralrs-core/src/cuda/hc_fused.rs new file mode 100644 index 000000000..c176e93ea --- /dev/null +++ b/mistralrs-core/src/cuda/hc_fused.rs @@ -0,0 +1,588 @@ +//! Parent system: ArcInfer / ArcMoE +//! +//! Rust side of `cuda/hc_fused.cu` — the fused DeepSeek-V4 router-region +//! kernels. +//! +//! Two entry points, each replacing a launch-bound candle op chain with one +//! kernel: +//! +//! - [`hc_pre_fused_cuda`] — the 18-launch RMS-scale + pre/post/comb scoring +//! block of [`crate::models::dsv4_mhc::V4MHCLayerParams::hc_pre`]. +//! - [`sqrt_softplus_cuda`] — the 9-launch `sqrt(softplus(x))` gate scoring +//! function of `MoeGate::forward`. +//! +//! Both are **bit-identical** to the chains they replace, not approximations. +//! The contract, and the three ways candle's arithmetic can be got wrong, are +//! documented at the top of `hc_fused.cu`. [`reference`] below carries scalar +//! replicas of both sides that `mod tests` asserts bit-identical without a GPU. +//! +//! Both are opt-out via `ARC_HC_FUSED=0`, which restores the eager chain. That +//! switch exists so the on-GPU A/B (fused vs. eager, greedy decode, compare +//! tokens) can be run from one binary. + +/// `ARC_HC_FUSED=0` disables the fused kernels and restores the eager candle +/// chains. Any other value (or unset) keeps them on. +/// +/// Read once and cached: this is consulted on every layer of every decode step, +/// and `std::env::var` takes a global lock. +pub fn fused_enabled() -> bool { + use std::sync::OnceLock; + static ENABLED: OnceLock = OnceLock::new(); + *ENABLED.get_or_init(|| !matches!(std::env::var("ARC_HC_FUSED").as_deref(), Ok("0"))) +} + +/// candle's `FastReduce` block size for a reduction of `len` elements +/// (`cuda_backend/mod.rs`: `usize::min(1024, el_to_sum_per_block) +/// .next_power_of_two()`). The fused kernel must use exactly this, because the +/// reduction tree's shape — and therefore the f32 rounding — depends on it. +pub(crate) fn candle_reduce_block_dim(len: usize) -> usize { + usize::min(1024, len).next_power_of_two() +} + +#[cfg(feature = "cuda")] +mod cuda_impl { + use candle_core as candle; + use candle_core::{DType, Result, Tensor}; + + use super::candle_reduce_block_dim; + + /// Pull a contiguous F32 CUDA tensor's device pointer. + fn f32_ptr(t: &Tensor, what: &str) -> Result<*const std::ffi::c_void> { + use candle_core::cuda_backend::cudarc::driver::DevicePtr; + if t.dtype() != DType::F32 { + candle::bail!("hc_fused: {what} must be F32, got {:?}", t.dtype()); + } + if !t.is_contiguous() { + candle::bail!("hc_fused: {what} must be contiguous"); + } + let (s, l) = t.storage_and_layout(); + let s = match &*s { + candle::Storage::Cuda(c) => c.as_cuda_slice::()?, + _ => candle::bail!("hc_fused: {what} must be on CUDA"), + }; + Ok(s.slice(l.start_offset()..).device_ptr(s.stream()).0 as *const std::ffi::c_void) + } + + /// Fused `hc_pre` middle section: the RMS statistic, its broadcast into + /// `mixes`, and the three scoring blocks. + /// + /// Inputs (all F32, contiguous, same CUDA device): + /// - `x_flat` `[n, d]` — the promoted residual stack, `d = hc * hidden` + /// - `mixes_raw` `[n, m]` — the gate GEMM output, `m = (2 + hc) * hc` + /// - `hc_scale` `[3]` + /// - `hc_base` `[m]` + /// + /// Returns `(pre [n, hc], post [n, hc], comb_pre [n, hc, hc])`. + #[allow(clippy::too_many_arguments)] + pub fn hc_pre_fused_cuda( + x_flat: &Tensor, + mixes_raw: &Tensor, + hc_scale: &Tensor, + hc_base: &Tensor, + hc: usize, + rms_eps: f64, + hc_eps: f64, + ) -> Result<(Tensor, Tensor, Tensor)> { + use candle_core::cuda_backend::cudarc::driver::DevicePtr; + + let (n, d) = x_flat.dims2()?; + let (n2, m) = mixes_raw.dims2()?; + if n != n2 { + candle::bail!("hc_fused: x_flat has {n} rows but mixes_raw has {n2}"); + } + if m != (2 + hc) * hc { + candle::bail!( + "hc_fused: mixes_raw has {m} columns, expected (2 + hc) * hc = {} for hc={hc}", + (2 + hc) * hc + ); + } + if hc_scale.dims1()? != 3 { + candle::bail!("hc_fused: hc_scale must be [3]"); + } + if hc_base.dims1()? != m { + candle::bail!("hc_fused: hc_base must be [{m}]"); + } + let block_dim = candle_reduce_block_dim(d); + if m > block_dim { + // The scoring tail is done by threads 0..m, so it must fit the + // block candle's reduction shape dictates. Unreachable for V4 + // (m = 24, block_dim = 1024) but a wrong answer if it ever isn't. + candle::bail!( + "hc_fused: m={m} exceeds the reduction block_dim={block_dim} implied by d={d}" + ); + } + if n == 0 { + candle::bail!("hc_fused: empty batch"); + } + + let x_ptr = f32_ptr(x_flat, "x_flat")?; + let mixes_ptr = f32_ptr(mixes_raw, "mixes_raw")?; + let scale_ptr = f32_ptr(hc_scale, "hc_scale")?; + let base_ptr = f32_ptr(hc_base, "hc_base")?; + + let dev = x_flat.device().as_cuda_device()?; + let pre_buf = unsafe { dev.alloc::(n * hc) }?; + let post_buf = unsafe { dev.alloc::(n * hc) }?; + let comb_buf = unsafe { dev.alloc::(n * hc * hc) }?; + let stream = dev.cuda_stream().cu_stream() as i64; + + #[allow(clippy::cast_possible_truncation)] + unsafe { + crate::cuda::ffi::hc_pre_fused_f32( + x_ptr, + mixes_ptr, + scale_ptr, + base_ptr, + pre_buf.device_ptr(pre_buf.stream()).0 as *mut std::ffi::c_void, + post_buf.device_ptr(post_buf.stream()).0 as *mut std::ffi::c_void, + comb_buf.device_ptr(comb_buf.stream()).0 as *mut std::ffi::c_void, + n as i32, + d as i32, + m as i32, + hc as i32, + block_dim as i32, + // candle's `mean_keepdim` scale is `T::from_f64(1f64 / len)`, + // i.e. the f64 reciprocal narrowed to f32 — not `1.0f32 / d`. + (1f64 / d as f64) as f32, + rms_eps as f32, + hc_eps as f32, + stream, + ); + } + + let wrap = |buf, shape| { + let st = candle::CudaStorage::wrap_cuda_slice(buf, dev.clone()); + Tensor::from((candle::Storage::Cuda(st), shape)) + }; + Ok(( + wrap(pre_buf, (n, hc)), + wrap(post_buf, (n, hc)), + wrap(comb_buf, (n, hc, hc)), + )) + } + + /// Fused `sqrt(softplus(x))` — the V4 gate scoring function. + pub fn sqrt_softplus_cuda(logits: &Tensor) -> Result { + use candle_core::cuda_backend::cudarc::driver::DevicePtr; + + let logits = logits.contiguous()?; + let numel = logits.elem_count(); + if numel == 0 { + candle::bail!("hc_fused: sqrt_softplus on an empty tensor"); + } + let in_ptr = f32_ptr(&logits, "logits")?; + let dev = logits.device().as_cuda_device()?; + let out_buf = unsafe { dev.alloc::(numel) }?; + let stream = dev.cuda_stream().cu_stream() as i64; + + #[allow(clippy::cast_possible_truncation)] + unsafe { + crate::cuda::ffi::sqrt_softplus_f32( + in_ptr, + out_buf.device_ptr(out_buf.stream()).0 as *mut std::ffi::c_void, + numel as i64, + stream, + ); + } + + let st = candle::CudaStorage::wrap_cuda_slice(out_buf, dev.clone()); + Ok(Tensor::from(( + candle::Storage::Cuda(st), + logits.shape().clone(), + ))) + } +} + +#[cfg(feature = "cuda")] +pub use cuda_impl::{hc_pre_fused_cuda, sqrt_softplus_cuda}; + +// Non-CUDA stubs so call sites need no `cfg`. They are never reached: every +// caller gates on [`usable`], which is false without a CUDA device. +#[cfg(not(feature = "cuda"))] +#[allow(clippy::too_many_arguments)] +pub fn hc_pre_fused_cuda( + _x_flat: &candle_core::Tensor, + _mixes_raw: &candle_core::Tensor, + _hc_scale: &candle_core::Tensor, + _hc_base: &candle_core::Tensor, + _hc: usize, + _rms_eps: f64, + _hc_eps: f64, +) -> candle_core::Result<( + candle_core::Tensor, + candle_core::Tensor, + candle_core::Tensor, +)> { + candle_core::bail!("hc_pre_fused_cuda requires the cuda feature") +} + +#[cfg(not(feature = "cuda"))] +pub fn sqrt_softplus_cuda(_logits: &candle_core::Tensor) -> candle_core::Result { + candle_core::bail!("sqrt_softplus_cuda requires the cuda feature") +} + +/// Whether the fused path should be taken for a tensor: a CUDA device, F32, +/// and not disabled by `ARC_HC_FUSED=0`. +/// +/// Checked BEFORE the call rather than by catching an error from it, so that a +/// genuine CUDA failure inside the kernel surfaces as an error instead of being +/// silently swallowed into the eager path (the "silent success" failure mode +/// this codebase has been bitten by before). +pub fn usable(t: &candle_core::Tensor) -> bool { + cfg!(feature = "cuda") + && t.device().is_cuda() + && t.dtype() == candle_core::DType::F32 + && fused_enabled() +} + +/// Scalar f32 replicas that pin the fused kernels' op ORDER and ROUNDING to +/// what the candle CUDA backend actually executes, bit for bit. +/// +/// Same discipline, and same known scalar-vs-GPU gap, as +/// [`crate::cuda::sinkhorn::reference`]: Rust/libm `f32::exp`/`ln` may differ +/// from CUDA libdevice `__nv_expf`/`__nv_logf` in the last ulp, but that +/// difference CANCELS in the on-GPU A/B because the candle chain and the fused +/// kernel call the same libdevice routine. Everything else here — add, mul, +/// div, fma, sqrt, single rounding per op — is exact IEEE f32 on both sides, so +/// the bitwise assertions in `mod tests` are meaningful. +#[allow(dead_code)] // consumed by `mod tests` +pub(crate) mod reference { + /// candle-kernels `cuda_utils.cuh`: `recipg(float a) { return 1.0 / a; }`. + /// The literal is a double, so this is `(float)(1.0 / (double)a)`. + pub(crate) fn candle_recip(a: f32) -> f32 { + (1.0f64 / a as f64) as f32 + } + + /// candle-kernels `unary.cu`: `sigmoid_fwd(x) = recipg(1 + expg(-x))`. + pub(crate) fn candle_sigmoid(x: f32) -> f32 { + candle_recip(1.0f32 + (-x).exp()) + } + + /// candle-kernels `fast_sum` (reduce.cu) over a contiguous row, with + /// candle's `FastReduce` block size. Each virtual thread accumulates its + /// strided slice sequentially into a zero-initialised accumulator, then the + /// pairwise tree runs over `block_dim` slots. + pub(crate) fn candle_fast_sum(vals: &[f32], block_dim: usize) -> f32 { + let mut shr = vec![0.0f32; block_dim]; + for (tid, slot) in shr.iter_mut().enumerate() { + let mut idx = tid; + while idx < vals.len() { + *slot += vals[idx]; + idx += block_dim; + } + } + let mut s = block_dim / 2; + while s > 0 { + for t in 0..s { + shr[t] += shr[t + s]; + } + s /= 2; + } + shr[0] + } + + /// The **candle op chain** of `dsv4_mhc::hc_pre`, one step per kernel + /// launch, in the order the eager path issues them. + /// + /// Returns `(pre, post, comb_pre)`. + #[allow(clippy::type_complexity)] + pub(crate) fn hc_pre_candle_replay( + x_flat: &[f32], + mixes_raw: &[f32], + hc_scale: &[f32], + hc_base: &[f32], + hc: usize, + rms_eps: f64, + hc_eps: f64, + ) -> (Vec, Vec, Vec) { + let d = x_flat.len(); + let m = mixes_raw.len(); + let block_dim = super::candle_reduce_block_dim(d); + + // sqr() -> a whole separate kernel, so x*x is rounded to f32 first. + let sq: Vec = x_flat.iter().map(|v| v * v).collect(); + // mean_keepdim = fast_sum, then a SEPARATE affine(1/d, 0) == fmaf. + let sum = candle_fast_sum(&sq, block_dim); + let inv_d = (1f64 / d as f64) as f32; + let mean = sum.mul_add(inv_d, 0.0f32); + // + rms_norm_eps -> affine(1.0, eps) == fmaf(x, 1, eps) + let meps = mean.mul_add(1.0f32, rms_eps as f32); + // recip() then sqrt() + let rsqrt = candle_recip(meps).sqrt(); + + // broadcast_mul -> `mixes` + let mixes: Vec = mixes_raw.iter().map(|v| v * rsqrt).collect(); + + let mut pre = Vec::with_capacity(hc); + for j in 0..hc { + // broadcast_mul, broadcast_add, sigmoid, affine(1, hc_eps) + let t = mixes[j] * hc_scale[0] + hc_base[j]; + pre.push(candle_sigmoid(t).mul_add(1.0f32, hc_eps as f32)); + } + let mut post = Vec::with_capacity(hc); + for j in 0..hc { + let t = mixes[hc + j] * hc_scale[1] + hc_base[hc + j]; + post.push(candle_sigmoid(t).mul_add(2.0f32, 0.0f32)); + } + let mut comb = Vec::with_capacity(hc * hc); + for j in 0..(m - 2 * hc) { + comb.push(mixes[2 * hc + j] * hc_scale[2] + hc_base[2 * hc + j]); + } + (pre, post, comb) + } + + /// The **fused kernel** of `hc_fused.cu`, transcribed: one block, `acc` in + /// a register, the same tree, then the scoring tail. + #[allow(clippy::type_complexity)] + pub(crate) fn hc_pre_fused_replay( + x_flat: &[f32], + mixes_raw: &[f32], + hc_scale: &[f32], + hc_base: &[f32], + hc: usize, + rms_eps: f64, + hc_eps: f64, + ) -> (Vec, Vec, Vec) { + let d = x_flat.len(); + let m = mixes_raw.len(); + let block_dim = super::candle_reduce_block_dim(d); + + // Per-thread: acc = fadd(acc, fmul(v, v)) over the strided slice. + let mut shr = vec![0.0f32; block_dim]; + for (tid, slot) in shr.iter_mut().enumerate() { + let mut acc = 0.0f32; + let mut idx = tid; + while idx < d { + let v = x_flat[idx]; + acc += v * v; + idx += block_dim; + } + *slot = acc; + } + let mut s = block_dim / 2; + while s > 0 { + for t in 0..s { + shr[t] += shr[t + s]; + } + s /= 2; + } + let sum = shr[0]; + let inv_d = (1f64 / d as f64) as f32; + let mean = sum.mul_add(inv_d, 0.0f32); + let meps = mean.mul_add(1.0f32, rms_eps as f32); + let rsqrt = candle_recip(meps).sqrt(); + + let mut pre = vec![0.0f32; hc]; + let mut post = vec![0.0f32; hc]; + let mut comb = vec![0.0f32; m - 2 * hc]; + for tid in 0..m { + let mx = mixes_raw[tid] * rsqrt; + if tid < hc { + let t = mx * hc_scale[0] + hc_base[tid]; + pre[tid] = candle_sigmoid(t).mul_add(1.0f32, hc_eps as f32); + } else if tid < 2 * hc { + let j = tid - hc; + let t = mx * hc_scale[1] + hc_base[hc + j]; + post[j] = candle_sigmoid(t).mul_add(2.0f32, 0.0f32); + } else { + let j = tid - 2 * hc; + let t = mx * hc_scale[2] + hc_base[2 * hc + j]; + comb[j] = t; + } + } + (pre, post, comb) + } + + /// The **candle op chain** of `MoeGate::forward`'s `SqrtSoftplus` arm. + pub(crate) fn sqrt_softplus_candle_replay(logits: &[f32]) -> Vec { + // max0 = maximum(logits, zeros_like) + let max0: Vec = logits.iter().map(|&l| l.max(0.0f32)).collect(); + // abs -> neg -> exp -> affine(1, 1) -> log + let a: Vec = logits.iter().map(|&l| l.abs()).collect(); + let n: Vec = a.iter().map(|&v| -v).collect(); + let e: Vec = n.iter().map(|&v| v.exp()).collect(); + let p1: Vec = e.iter().map(|&v| v.mul_add(1.0f32, 1.0f32)).collect(); + let lg: Vec = p1.iter().map(|&v| v.ln()).collect(); + // badd then usqrt + max0.iter() + .zip(lg.iter()) + .map(|(&x, &y)| (x + y).sqrt()) + .collect() + } + + /// The **fused kernel** of `hc_fused.cu`, transcribed. + pub(crate) fn sqrt_softplus_fused_replay(logits: &[f32]) -> Vec { + logits + .iter() + .map(|&l| { + let mx = l.max(0.0f32); + let a = l.abs(); + let e = (-a).exp(); + let p1 = e.mul_add(1.0f32, 1.0f32); + let lg = p1.ln(); + (mx + lg).sqrt() + }) + .collect() + } +} + +#[cfg(test)] +mod tests { + use super::reference::*; + + /// Deterministic LCG so a failure is reproducible without a fixture file. + struct Lcg(u64); + impl Lcg { + fn next_f32(&mut self, lo: f32, hi: f32) -> f32 { + self.0 = self.0.wrapping_mul(6364136223846793005).wrapping_add(1442695040888963407); + let u = ((self.0 >> 40) as f32) / ((1u32 << 24) as f32); + lo + u * (hi - lo) + } + } + + #[test] + fn candle_reduce_block_dim_matches_fast_reduce() { + // cuda_backend/mod.rs: usize::min(1024, el_to_sum_per_block).next_power_of_two() + assert_eq!(super::candle_reduce_block_dim(16384), 1024); // V4: hc_mult * hidden + assert_eq!(super::candle_reduce_block_dim(1024), 1024); + assert_eq!(super::candle_reduce_block_dim(1000), 1024); + assert_eq!(super::candle_reduce_block_dim(24), 32); + assert_eq!(super::candle_reduce_block_dim(4), 4); + assert_eq!(super::candle_reduce_block_dim(3), 4); + } + + /// The whole point of the file: the fused kernel and the candle op chain + /// must agree BITWISE, at V4's real shapes. + #[test] + fn hc_pre_fused_is_bit_identical_to_candle_chain() { + let hc = 4usize; + let hidden = 4096usize; + let d = hc * hidden; + let m = (2 + hc) * hc; + let mut rng = Lcg(0x5eed_1234_abcd_0001); + + for trial in 0..8 { + let x: Vec = (0..d).map(|_| rng.next_f32(-3.0, 3.0)).collect(); + let mixes: Vec = (0..m).map(|_| rng.next_f32(-6.0, 6.0)).collect(); + let scale: Vec = (0..3).map(|_| rng.next_f32(-2.0, 2.0)).collect(); + let base: Vec = (0..m).map(|_| rng.next_f32(-2.0, 2.0)).collect(); + + let (pa, qa, ca) = hc_pre_candle_replay(&x, &mixes, &scale, &base, hc, 1e-6, 1e-6); + let (pb, qb, cb) = hc_pre_fused_replay(&x, &mixes, &scale, &base, hc, 1e-6, 1e-6); + + for (i, (a, b)) in pa.iter().zip(pb.iter()).enumerate() { + assert_eq!(a.to_bits(), b.to_bits(), "trial {trial} pre[{i}]: {a} vs {b}"); + } + for (i, (a, b)) in qa.iter().zip(qb.iter()).enumerate() { + assert_eq!(a.to_bits(), b.to_bits(), "trial {trial} post[{i}]: {a} vs {b}"); + } + for (i, (a, b)) in ca.iter().zip(cb.iter()).enumerate() { + assert_eq!(a.to_bits(), b.to_bits(), "trial {trial} comb[{i}]: {a} vs {b}"); + } + } + } + + /// A guard that cannot go red is not a guard. Reassociating the reduction — + /// the single most likely way to break this fusion — must be detected. + #[test] + fn hc_pre_guard_detects_a_reassociated_reduction() { + let hc = 4usize; + let d = hc * 4096; + let m = (2 + hc) * hc; + let mut rng = Lcg(0xdead_beef_0000_0007); + let x: Vec = (0..d).map(|_| rng.next_f32(-3.0, 3.0)).collect(); + let mixes: Vec = (0..m).map(|_| rng.next_f32(-6.0, 6.0)).collect(); + let scale: Vec = (0..3).map(|_| rng.next_f32(-2.0, 2.0)).collect(); + let base: Vec = (0..m).map(|_| rng.next_f32(-2.0, 2.0)).collect(); + + let (pre_ref, ..) = hc_pre_candle_replay(&x, &mixes, &scale, &base, hc, 1e-6, 1e-6); + + // Naive left-to-right sum instead of candle's strided-then-tree order. + let mut seq = 0.0f32; + for v in &x { + seq += v * v; + } + let inv_d = (1f64 / d as f64) as f32; + let rsqrt = candle_recip(seq.mul_add(inv_d, 0.0f32).mul_add(1.0f32, 1e-6f32)).sqrt(); + let pre_wrong: Vec = (0..hc) + .map(|j| { + let t = (mixes[j] * rsqrt) * scale[0] + base[j]; + candle_sigmoid(t).mul_add(1.0f32, 1e-6f32) + }) + .collect(); + + assert!( + pre_ref + .iter() + .zip(pre_wrong.iter()) + .any(|(a, b)| a.to_bits() != b.to_bits()), + "the reduction-order guard is vacuous: a sequential sum of {d} squares produced \ + bit-identical output to candle's strided+tree order, so this test would pass on a \ + kernel that reassociates the reduction" + ); + } + + #[test] + fn sqrt_softplus_fused_is_bit_identical_to_candle_chain() { + let mut rng = Lcg(0x0bad_c0de_1111_2222); + // n_routed_experts = 256 for V4-Flash. + let logits: Vec = (0..256).map(|_| rng.next_f32(-40.0, 40.0)).collect(); + let a = sqrt_softplus_candle_replay(&logits); + let b = sqrt_softplus_fused_replay(&logits); + for (i, (x, y)) in a.iter().zip(b.iter()).enumerate() { + assert_eq!(x.to_bits(), y.to_bits(), "logit[{i}]={}: {x} vs {y}", logits[i]); + } + // Including the exact-zero / sign boundary that `maximum(x, 0)` and + // `abs(x)` disagree on. + let edge = [0.0f32, -0.0f32, 1e-30, -1e-30, 88.0, -88.0, 1.0, -1.0]; + let a = sqrt_softplus_candle_replay(&edge); + let b = sqrt_softplus_fused_replay(&edge); + for (i, (x, y)) in a.iter().zip(b.iter()).enumerate() { + assert_eq!(x.to_bits(), y.to_bits(), "edge[{i}]={}: {x} vs {y}", edge[i]); + } + } + + /// Same discipline for the softplus guard: prove it can go red. + /// + /// Mutation-tests the three ways a careless reimplementation of + /// `sqrt(softplus(x))` actually goes wrong, and asserts the bitwise + /// comparison catches each. The first one is the reason this test exists at + /// all: an earlier version of it sampled only `[-30, -12.5, -1, 0.5, 3.25, + /// 17]` and PASSED against the naive form, because over that range the two + /// formulations happen to agree bit for bit — the stable split only earns + /// its keep once `exp(x)` overflows f32 near x = 88.7. A guard that green- + /// lights the unstable form is not a guard. + #[test] + fn sqrt_softplus_guard_detects_plausible_mistranscriptions() { + // Spans the overflow boundary (exp overflows f32 above ~88.7) and the + // sign boundary that `max(x, 0)` and `-|x|` turn on. + let logits: Vec = vec![-120.0, -30.0, -1.0, -0.25, 0.0, 0.25, 3.25, 17.0, 95.0, 200.0]; + let reference = sqrt_softplus_candle_replay(&logits); + + let mutants: [(&str, fn(f32) -> f32); 3] = [ + // 1. the unstable form: log(1 + exp(x)), no max/abs split. + ("unstable log(1+exp(x))", |l| (1.0f32 + l.exp()).ln().sqrt()), + // 2. dropped the max(x, 0) term. + ("dropped max(x,0)", |l| { + (1.0f32 + (-l.abs()).exp()).ln().sqrt() + }), + // 3. lost the negation on the exponent. + ("exp(+|x|) instead of exp(-|x|)", |l| { + (l.max(0.0f32) + (1.0f32 + l.abs().exp()).ln()).sqrt() + }), + ]; + + for (name, mutate) in mutants { + let got: Vec = logits.iter().map(|&l| mutate(l)).collect(); + assert!( + reference + .iter() + .zip(got.iter()) + .any(|(a, b)| a.to_bits() != b.to_bits()), + "the softplus guard is vacuous for mutation '{name}': it produced bit-identical \ + output to the reference on every sample, so this test would pass on a kernel \ + carrying that bug" + ); + } + } +} diff --git a/mistralrs-core/src/cuda/mod.rs b/mistralrs-core/src/cuda/mod.rs index 7318477f8..68f0e65ff 100644 --- a/mistralrs-core/src/cuda/mod.rs +++ b/mistralrs-core/src/cuda/mod.rs @@ -1,5 +1,6 @@ pub mod ffi; pub mod gdn; +pub mod hc_fused; pub mod moe; pub mod sinkhorn; pub mod ssm; diff --git a/mistralrs-core/src/cuda/sinkhorn.cu b/mistralrs-core/src/cuda/sinkhorn.cu index ae665e64b..767ac55a1 100644 --- a/mistralrs-core/src/cuda/sinkhorn.cu +++ b/mistralrs-core/src/cuda/sinkhorn.cu @@ -70,123 +70,195 @@ // hc_mult is 4 for V4-Flash; cap at 16 for safety (shared-mem sized at launch). #define SINKHORN_MAX_HC 16 +// --------------------------------------------------------------------------- +// PERFORMANCE NOTE (why this file is templated on `hc`) +// --------------------------------------------------------------------------- +// The first version of this kernel took `hc` as a RUNTIME argument and held the +// per-thread row / tree buffers in `float r[SINKHORN_MAX_HC]` arrays walked by +// runtime-bounded loops. nvcc cannot keep an array in registers unless every +// index is resolvable at compile time, so all three arrays were demoted to +// LOCAL MEMORY -- `ptxas -v` reported `192 bytes stack frame` (48 floats = +// r[16] + buf[16] + col[16]). With hc = 4 the kernel launches ONE block of FOUR +// threads, so there is no occupancy to hide that latency behind: every one of +// the ~1,800 dependent local-memory round trips per call (20 iterations x two +// tree reductions x two arrays, plus the row read/write) is exposed. Measured +// cost: 30.0 us per call on a [1, 4, 4] tensor -- 2.549 ms/step over the 86 +// calls, 28% of the whole router region's GPU time, for 16 floats of work. +// +// Templating on `hc` makes every index a compile-time constant, the arrays +// become registers (`0 bytes stack frame`), and the arithmetic is UNCHANGED -- +// same ops, same order, same intrinsics -- so bit-identity is preserved by +// construction. The tree reductions are expressed as template recursion rather +// than `#pragma unroll` loops specifically so that the halving step `S` is a +// constant in the type system and cannot silently fall back to a runtime index. +// --------------------------------------------------------------------------- + namespace { -__device__ __forceinline__ int next_pow2_le16(int v) { +constexpr int next_pow2_ce(int v) { int p = 1; while (p < v) p <<= 1; return p; } +// One level of candle's pairwise tree, with the stride `S` fixed by the type. +// `buf` is taken by array reference (not pointer) so the indices stay +// compile-time and the storage stays in registers. +template struct TreeSumLevel { + __device__ __forceinline__ static void run(float (&buf)[N]) { + constexpr int S = P / 2; +#pragma unroll + for (int t = 0; t < S; ++t) { + buf[t] = __fadd_rn(buf[t], buf[t + S]); + } + TreeSumLevel::run(buf); + } +}; +template struct TreeSumLevel<1, N> { + __device__ __forceinline__ static void run(float (&)[N]) {} +}; + +template struct TreeMaxLevel { + __device__ __forceinline__ static void run(float (&buf)[N]) { + constexpr int S = P / 2; +#pragma unroll + for (int t = 0; t < S; ++t) { + buf[t] = fmaxf(buf[t], buf[t + S]); + } + TreeMaxLevel::run(buf); + } +}; +template struct TreeMaxLevel<1, N> { + __device__ __forceinline__ static void run(float (&)[N]) {} +}; + // Replays candle-kernels `fast_sum` (reduce.cu) for reduced_len <= 16: // zero-initialized accumulators, one element per virtual thread // (block_dim = next_pow2(len) >= len so each thread loads at most one), // then the pairwise tree. `__fadd_rn(0.0f, v)` mirrors `shr[tid] = 0; // shr[tid] += v` (note: turns -0.0f into +0.0f, exactly like candle). -__device__ __forceinline__ float candle_tree_sum(const float* v, int len) { - float buf[SINKHORN_MAX_HC]; - const int p = next_pow2_le16(len); - for (int t = 0; t < p; ++t) { - buf[t] = (t < len) ? __fadd_rn(0.0f, v[t]) : 0.0f; - } - for (int s = p >> 1; s > 0; s >>= 1) { - for (int t = 0; t < s; ++t) { - buf[t] = __fadd_rn(buf[t], buf[t + s]); - } +template __device__ __forceinline__ float candle_tree_sum(const float (&v)[HC]) { + constexpr int P = next_pow2_ce(HC); + float buf[P]; +#pragma unroll + for (int t = 0; t < P; ++t) { + buf[t] = (t < HC) ? __fadd_rn(0.0f, v[t]) : 0.0f; } + TreeSumLevel::run(buf); return buf[0]; } // Replays candle-kernels `fast_max` (reduce.cu): -INF init, maxg == fmaxf // (NaN-ignoring IEEE maxNum), pairwise tree. Order-insensitive for finite // inputs but mirrored anyway so NaN propagation matches candle exactly. -__device__ __forceinline__ float candle_tree_max(const float* v, int len) { - float buf[SINKHORN_MAX_HC]; - const int p = next_pow2_le16(len); - for (int t = 0; t < p; ++t) { - buf[t] = (t < len) ? fmaxf(-INFINITY, v[t]) : -INFINITY; - } - for (int s = p >> 1; s > 0; s >>= 1) { - for (int t = 0; t < s; ++t) { - buf[t] = fmaxf(buf[t], buf[t + s]); - } +template __device__ __forceinline__ float candle_tree_max(const float (&v)[HC]) { + constexpr int P = next_pow2_ce(HC); + float buf[P]; +#pragma unroll + for (int t = 0; t < P; ++t) { + buf[t] = (t < HC) ? fmaxf(-INFINITY, v[t]) : -INFINITY; } + TreeMaxLevel::run(buf); return buf[0]; } -} // namespace - -extern "C" { - // in/out: [n, hc, hc] row-major F32. One block per matrix `n`, `hc` threads. -// Shared memory: hc*hc (matrix tile) + hc (column sums) floats. +// `hc` is a template parameter so the per-thread row / tree buffers stay in +// registers; see the PERFORMANCE NOTE above. The arithmetic is identical to the +// runtime-`hc` version this replaced. The kernel lives inside the anonymous +// namespace because a template cannot have C linkage; only the dispatcher below +// is `extern "C"`, which is all the Rust FFI binds to. +template __global__ void sinkhorn_normalize_f32_kernel( const float* __restrict__ in, float* __restrict__ out, int n, - int hc, int iters, float eps ) { int batch = blockIdx.x; int row = threadIdx.x; - if (batch >= n || row >= hc) return; + if (batch >= n || row >= HC) return; - extern __shared__ float smem[]; - float* mat = smem; // [hc * hc] - float* csum = smem + hc * hc; // [hc] + __shared__ float mat[HC * HC]; + __shared__ float csum[HC]; - const float* my_in = in + (size_t)batch * hc * hc + (size_t)row * hc; + const float* my_in = in + (size_t)batch * HC * HC + (size_t)row * HC; // Each thread owns one row in registers. - float r[SINKHORN_MAX_HC]; - for (int j = 0; j < hc; ++j) r[j] = my_in[j]; + float r[HC]; +#pragma unroll + for (int j = 0; j < HC; ++j) r[j] = my_in[j]; // ---- 1. stable row softmax, then + eps ---- // candle: max_keepdim(-1) -> broadcast_sub -> exp -> sum_keepdim(-1) // -> broadcast_div -> affine(+eps) - const float m = candle_tree_max(r, hc); - for (int j = 0; j < hc; ++j) r[j] = expf(__fsub_rn(r[j], m)); - const float rs = candle_tree_sum(r, hc); - for (int j = 0; j < hc; ++j) r[j] = __fadd_rn(__fdiv_rn(r[j], rs), eps); + const float m = candle_tree_max(r); +#pragma unroll + for (int j = 0; j < HC; ++j) r[j] = expf(__fsub_rn(r[j], m)); + const float rs = candle_tree_sum(r); +#pragma unroll + for (int j = 0; j < HC; ++j) r[j] = __fadd_rn(__fdiv_rn(r[j], rs), eps); // publish row to shared - for (int j = 0; j < hc; ++j) mat[row * hc + j] = r[j]; +#pragma unroll + for (int j = 0; j < HC; ++j) mat[row * HC + j] = r[j]; __syncthreads(); // ---- 2. initial column normalize: x / (colsum + eps) ---- // column `row` sum = tree-sum over rows k of mat[k][row] { - float col[SINKHORN_MAX_HC]; - for (int k = 0; k < hc; ++k) col[k] = mat[k * hc + row]; - csum[row] = __fadd_rn(candle_tree_sum(col, hc), eps); + float col[HC]; +#pragma unroll + for (int k = 0; k < HC; ++k) col[k] = mat[k * HC + row]; + csum[row] = __fadd_rn(candle_tree_sum(col), eps); } __syncthreads(); - for (int j = 0; j < hc; ++j) r[j] = __fdiv_rn(mat[row * hc + j], csum[j]); +#pragma unroll + for (int j = 0; j < HC; ++j) r[j] = __fdiv_rn(mat[row * HC + j], csum[j]); // ---- 3. (iters - 1) more row->col passes ---- for (int it = 0; it < iters - 1; ++it) { // row normalize: x / (rowsum + eps) (r holds this thread's row) - const float rsum = __fadd_rn(candle_tree_sum(r, hc), eps); - for (int j = 0; j < hc; ++j) r[j] = __fdiv_rn(r[j], rsum); - for (int j = 0; j < hc; ++j) mat[row * hc + j] = r[j]; + const float rsum = __fadd_rn(candle_tree_sum(r), eps); +#pragma unroll + for (int j = 0; j < HC; ++j) r[j] = __fdiv_rn(r[j], rsum); +#pragma unroll + for (int j = 0; j < HC; ++j) mat[row * HC + j] = r[j]; __syncthreads(); // column normalize: x / (colsum + eps) { - float col[SINKHORN_MAX_HC]; - for (int k = 0; k < hc; ++k) col[k] = mat[k * hc + row]; - csum[row] = __fadd_rn(candle_tree_sum(col, hc), eps); + float col[HC]; +#pragma unroll + for (int k = 0; k < HC; ++k) col[k] = mat[k * HC + row]; + csum[row] = __fadd_rn(candle_tree_sum(col), eps); } __syncthreads(); - for (int j = 0; j < hc; ++j) r[j] = __fdiv_rn(mat[row * hc + j], csum[j]); +#pragma unroll + for (int j = 0; j < HC; ++j) r[j] = __fdiv_rn(mat[row * HC + j], csum[j]); } // ---- write out ---- - float* my_out = out + (size_t)batch * hc * hc + (size_t)row * hc; - for (int j = 0; j < hc; ++j) my_out[j] = r[j]; + float* my_out = out + (size_t)batch * HC * HC + (size_t)row * HC; +#pragma unroll + for (int j = 0; j < HC; ++j) my_out[j] = r[j]; } +// Every hc in [1, SINKHORN_MAX_HC] gets its own instantiation, so there is no +// runtime-`hc` fallback whose numerics could drift from the templated path. +// V4-Flash uses hc = 4; the rest are cheap (a few hundred bytes of cubin each) +// and keep the Rust-side contract "hc <= 16 is supported" honest. +#define SINKHORN_LAUNCH(HC_) \ + case HC_: \ + sinkhorn_normalize_f32_kernel<<>>( \ + (const float*)in, (float*)out, n, iters, eps); \ + return; + +} // namespace + +extern "C" { + void sinkhorn_normalize_f32( const void* in, void* out, @@ -198,9 +270,30 @@ void sinkhorn_normalize_f32( ) { dim3 grid(n, 1, 1); dim3 block(hc, 1, 1); - size_t shmem = (size_t)(hc * hc + hc) * sizeof(float); - sinkhorn_normalize_f32_kernel<<>>( - (const float*)in, (float*)out, n, hc, iters, eps); + switch (hc) { + SINKHORN_LAUNCH(1) + SINKHORN_LAUNCH(2) + SINKHORN_LAUNCH(3) + SINKHORN_LAUNCH(4) + SINKHORN_LAUNCH(5) + SINKHORN_LAUNCH(6) + SINKHORN_LAUNCH(7) + SINKHORN_LAUNCH(8) + SINKHORN_LAUNCH(9) + SINKHORN_LAUNCH(10) + SINKHORN_LAUNCH(11) + SINKHORN_LAUNCH(12) + SINKHORN_LAUNCH(13) + SINKHORN_LAUNCH(14) + SINKHORN_LAUNCH(15) + SINKHORN_LAUNCH(16) + default: + // Unreachable: sinkhorn_normalize_cuda rejects hc > SINKHORN_MAX_HC + // before calling. Leaving `out` untouched here would be a silent wrong + // answer, so do nothing and let the caller's guard be the contract. + return; + } } +#undef SINKHORN_LAUNCH } // extern "C" diff --git a/mistralrs-core/src/models/deepseek4.rs b/mistralrs-core/src/models/deepseek4.rs index e03103a1f..f18e6a093 100644 --- a/mistralrs-core/src/models/deepseek4.rs +++ b/mistralrs-core/src/models/deepseek4.rs @@ -2085,11 +2085,23 @@ impl MoeGate { // V4: sqrt(softplus(x)). Stable formulation: // softplus(x) = max(x, 0) + log(1 + exp(-|x|)). // Audit §8 P1 item 14. + // + // The eager form below is NINE kernel launches (`zeros_like`, + // `bmaximum`, `uabs`, `uneg`, `uexp`, `affine`, `ulog`, `badd`, + // `usqrt`) on a `[1, n_routed_experts]` = [1, 256] tensor, once per + // MoE layer per token. `cuda/hc_fused.cu` collapses it to one, + // bit-identically — this expression decides WHICH EXPERTS RUN, so + // the fused kernel transcribes candle's ops rather than + // re-deriving them. `ARC_HC_FUSED=0` restores the chain for A/B. ScoringFunc::SqrtSoftplus => { - let max0 = logits.maximum(&logits.zeros_like()?)?; - let abs = logits.abs()?; - let softplus = (max0 + ((abs.neg()?.exp()? + 1.0)?.log()?))?; - softplus.sqrt()? + if crate::cuda::hc_fused::usable(&logits) { + crate::cuda::hc_fused::sqrt_softplus_cuda(&logits)? + } else { + let max0 = logits.maximum(&logits.zeros_like()?)?; + let abs = logits.abs()?; + let softplus = (max0 + ((abs.neg()?.exp()? + 1.0)?.log()?))?; + softplus.sqrt()? + } } }; drop(_prof_score); diff --git a/mistralrs-core/src/models/dsv4_mhc.rs b/mistralrs-core/src/models/dsv4_mhc.rs index 00afc5940..d3c740ed7 100644 --- a/mistralrs-core/src/models/dsv4_mhc.rs +++ b/mistralrs-core/src/models/dsv4_mhc.rs @@ -278,57 +278,109 @@ impl V4MHCLayerParams { // Promote to F32, flatten to [N, hc*h]. let x_flat = x.reshape((n, hc * h))?.to_dtype(DType::F32)?; - // rsqrt(mean(x^2) + eps) - let sq_mean = x_flat.sqr()?.mean_keepdim(D::Minus1)?; - let rsqrt = (sq_mean + self.rt.rms_norm_eps)?.recip()?.sqrt()?; // [N, 1] - - // mixes = (x_flat @ fn^T) * rsqrt → [N, mix_hc] // Defensively cast weight tensors to F32 — try_load already produces F32, // but hand-constructed callers (tests, external integrators) may not. let hc_fn_f32 = hc_fn.to_dtype(DType::F32)?; let hc_scale_f32 = hc_scale.to_dtype(DType::F32)?; let hc_base_f32 = hc_base.to_dtype(DType::F32)?; let mixes_raw = x_flat.matmul(&hc_fn_f32.t()?)?; - let mixes = mixes_raw.broadcast_mul(&rsqrt)?; - // Slot indices in `mixes`: - // pre : [.., 0 .. hc) - // post : [.., hc .. 2*hc) - // comb : [.., 2*hc .. (2+hc)*hc) reshape to [.., hc, hc] - let pre_block = mixes.narrow(D::Minus1, 0, hc)?; - let post_block = mixes.narrow(D::Minus1, hc, hc)?; - let comb_block = mixes - .narrow(D::Minus1, 2 * hc, hc * hc)? - .reshape((n, hc, hc))?; - - // hc_scale is [3]; hc_base is [mix_hc] split into three blocks. - let s_pre = hc_scale_f32.narrow(0, 0, 1)?; - let s_post = hc_scale_f32.narrow(0, 1, 1)?; - let s_comb = hc_scale_f32.narrow(0, 2, 1)?; - let b_pre = hc_base_f32.narrow(0, 0, hc)?; - let b_post = hc_base_f32.narrow(0, hc, hc)?; - let b_comb = hc_base_f32.narrow(0, 2 * hc, hc * hc)?.reshape((hc, hc))?; - - // pre = sigmoid(pre_block * s_pre + b_pre) + eps - let pre = candle_nn::ops::sigmoid( - &(pre_block.broadcast_mul(&s_pre)?.broadcast_add(&b_pre)?), - )?; - let pre = (pre + self.rt.hc_eps)?; - - // post = 2 * sigmoid(post_block * s_post + b_post) - // NOTE: use affine() for the scalar *2 rather than a device-scalar - // Tensor::new(2f32, device) — the latter is a per-call CPU->GPU sync - // (CLAUDE.md pitfall #5) that breaks CUDA-graph capture of the decode - // forward. affine folds the constant into the kernel, no allocation. - let post_sig = candle_nn::ops::sigmoid( - &(post_block.broadcast_mul(&s_post)?.broadcast_add(&b_post)?), - )?; - let post = post_sig.affine(2.0, 0.0)?; + // Everything from here to `comb_pre` is ONE fused kernel on CUDA. The + // eager chain below it is 18 launches — a 7-launch hand-decomposed RMS + // statistic (`sqr -> fast_sum -> affine -> affine -> recip -> sqrt -> + // bmul`) plus 11 for the three scoring blocks — all on 24 floats once + // the reduction is done, twice per layer, 43 layers. At b=1 that is + // ~1,460 of the step's ~7,900 kernel launches, i.e. pure launch + // overhead. `cuda/hc_fused.cu` is bit-identical to this chain by + // construction, not by tolerance; the eager path stays reachable via + // `ARC_HC_FUSED=0` so the two can be A/B'd from one binary. + let on_cuda = crate::cuda::hc_fused::usable(&x_flat); + let shapes_ok = x_flat.is_contiguous() + && mixes_raw.is_contiguous() + && hc_scale_f32.is_contiguous() + && hc_base_f32.is_contiguous() + && hc_base_f32.dims1().map(|v| v == (2 + hc) * hc).unwrap_or(false) + && hc_scale_f32.dims1().map(|v| v == 3).unwrap_or(false); + if on_cuda && !shapes_ok { + // Falling back on CUDA is a silent 18-launch regression that looks + // exactly like "the fusion didn't help". Say so once rather than + // letting a layout change quietly undo the optimisation. + static WARNED: std::sync::Once = std::sync::Once::new(); + WARNED.call_once(|| { + tracing::warn!( + "V4 mHC: fused hc_pre kernel unusable (x_flat contig={}, mixes contig={}, \ + scale contig={} dims={:?}, base contig={} dims={:?}) — falling back to the \ + 18-launch eager chain.", + x_flat.is_contiguous(), + mixes_raw.is_contiguous(), + hc_scale_f32.is_contiguous(), + hc_scale_f32.dims(), + hc_base_f32.is_contiguous(), + hc_base_f32.dims(), + ); + }); + } + let fused = on_cuda && shapes_ok; + + let (pre, post, comb_pre) = if fused { + crate::cuda::hc_fused::hc_pre_fused_cuda( + &x_flat, + &mixes_raw, + &hc_scale_f32, + &hc_base_f32, + hc, + self.rt.rms_norm_eps, + self.rt.hc_eps, + )? + } else { + // rsqrt(mean(x^2) + eps) + let sq_mean = x_flat.sqr()?.mean_keepdim(D::Minus1)?; + let rsqrt = (sq_mean + self.rt.rms_norm_eps)?.recip()?.sqrt()?; // [N, 1] + + // mixes = (x_flat @ fn^T) * rsqrt → [N, mix_hc] + let mixes = mixes_raw.broadcast_mul(&rsqrt)?; + + // Slot indices in `mixes`: + // pre : [.., 0 .. hc) + // post : [.., hc .. 2*hc) + // comb : [.., 2*hc .. (2+hc)*hc) reshape to [.., hc, hc] + let pre_block = mixes.narrow(D::Minus1, 0, hc)?; + let post_block = mixes.narrow(D::Minus1, hc, hc)?; + let comb_block = mixes + .narrow(D::Minus1, 2 * hc, hc * hc)? + .reshape((n, hc, hc))?; + + // hc_scale is [3]; hc_base is [mix_hc] split into three blocks. + let s_pre = hc_scale_f32.narrow(0, 0, 1)?; + let s_post = hc_scale_f32.narrow(0, 1, 1)?; + let s_comb = hc_scale_f32.narrow(0, 2, 1)?; + let b_pre = hc_base_f32.narrow(0, 0, hc)?; + let b_post = hc_base_f32.narrow(0, hc, hc)?; + let b_comb = hc_base_f32.narrow(0, 2 * hc, hc * hc)?.reshape((hc, hc))?; + + // pre = sigmoid(pre_block * s_pre + b_pre) + eps + let pre = candle_nn::ops::sigmoid( + &(pre_block.broadcast_mul(&s_pre)?.broadcast_add(&b_pre)?), + )?; + let pre = (pre + self.rt.hc_eps)?; + + // post = 2 * sigmoid(post_block * s_post + b_post) + // NOTE: use affine() for the scalar *2 rather than a device-scalar + // Tensor::new(2f32, device) — the latter is a per-call CPU->GPU sync + // (CLAUDE.md pitfall #5) that breaks CUDA-graph capture of the decode + // forward. affine folds the constant into the kernel, no allocation. + let post_sig = candle_nn::ops::sigmoid( + &(post_block.broadcast_mul(&s_post)?.broadcast_add(&b_post)?), + )?; + let post = post_sig.affine(2.0, 0.0)?; + + let comb_pre = comb_block + .broadcast_mul(&s_comb)? + .broadcast_add(&b_comb)?; // [N, hc, hc] + (pre, post, comb_pre) + }; // comb = sinkhorn_normalize(comb_block * s_comb + b_comb) - let comb_pre = comb_block - .broadcast_mul(&s_comb)? - .broadcast_add(&b_comb)?; // [N, hc, hc] let comb = sinkhorn_normalize(&comb_pre, self.rt.hc_sinkhorn_iters, self.rt.hc_eps)?; // y = sum_i pre[..., i, None] * x[..., i, :] → [N, hidden] From a92cf6d11f720480e38bc5a47fe0f007fb93cb58 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick Date: Wed, 19 Aug 2026 22:05:45 +0100 Subject: [PATCH 20/22] =?UTF-8?q?fix(ArcMoE):=20cuda-gated=20build=20?= =?UTF-8?q?=E2=80=94=20device=5Fptr=20guard=20lifetime=20and=20monomorphic?= =?UTF-8?q?=20wrap=20closure?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01SpVNMpb13HkUXqSqbN1o9H --- mistralrs-core/src/cuda/hc_fused.rs | 20 ++++++++++++-------- 1 file changed, 12 insertions(+), 8 deletions(-) diff --git a/mistralrs-core/src/cuda/hc_fused.rs b/mistralrs-core/src/cuda/hc_fused.rs index c176e93ea..e0ad69425 100644 --- a/mistralrs-core/src/cuda/hc_fused.rs +++ b/mistralrs-core/src/cuda/hc_fused.rs @@ -60,7 +60,12 @@ mod cuda_impl { candle::Storage::Cuda(c) => c.as_cuda_slice::()?, _ => candle::bail!("hc_fused: {what} must be on CUDA"), }; - Ok(s.slice(l.start_offset()..).device_ptr(s.stream()).0 as *const std::ffi::c_void) + // Bind before returning: `device_ptr` hands back a (ptr, guard) pair + // whose guard borrows `s`, so the pointer must be extracted in its own + // statement rather than in tail position. Same shape as + // `sinkhorn::sinkhorn_normalize_cuda`. + let ptr = s.slice(l.start_offset()..).device_ptr(s.stream()).0 as *const std::ffi::c_void; + Ok(ptr) } /// Fused `hc_pre` middle section: the RMS statistic, its broadcast into @@ -150,14 +155,13 @@ mod cuda_impl { ); } - let wrap = |buf, shape| { - let st = candle::CudaStorage::wrap_cuda_slice(buf, dev.clone()); - Tensor::from((candle::Storage::Cuda(st), shape)) - }; + let pre_st = candle::CudaStorage::wrap_cuda_slice(pre_buf, dev.clone()); + let post_st = candle::CudaStorage::wrap_cuda_slice(post_buf, dev.clone()); + let comb_st = candle::CudaStorage::wrap_cuda_slice(comb_buf, dev.clone()); Ok(( - wrap(pre_buf, (n, hc)), - wrap(post_buf, (n, hc)), - wrap(comb_buf, (n, hc, hc)), + Tensor::from((candle::Storage::Cuda(pre_st), (n, hc))), + Tensor::from((candle::Storage::Cuda(post_st), (n, hc))), + Tensor::from((candle::Storage::Cuda(comb_st), (n, hc, hc))), )) } From 48ae6b70ca9aa66404ec2ba6d3846a54350a4825 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick Date: Wed, 19 Aug 2026 22:13:16 +0100 Subject: [PATCH 21/22] test(ArcMoE): extend the build-wiring tripwire to hc_fused.cu MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The sinkhorn guard asserted build.rs contained the literal `.exclude(&["sinkhorn.cu"])`, so adding a second bit-identity-critical kernel to that list failed it for the one reason that is not a regression — the kind of failure that gets 'fixed' by deleting the assertion. It now matches on the exclude list's CONTENTS, and hc_fused.rs carries the mirror guard: fast-math #error present, IEEE intrinsics present, __expf/__logf/ __fdividef/rsqrtf absent, and the file wired into the --fmad=false builder. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01SpVNMpb13HkUXqSqbN1o9H --- mistralrs-core/src/cuda/hc_fused.rs | 46 +++++++++++++++++++++++++++++ mistralrs-core/src/cuda/sinkhorn.rs | 18 +++++++++-- 2 files changed, 62 insertions(+), 2 deletions(-) diff --git a/mistralrs-core/src/cuda/hc_fused.rs b/mistralrs-core/src/cuda/hc_fused.rs index e0ad69425..7376ce005 100644 --- a/mistralrs-core/src/cuda/hc_fused.rs +++ b/mistralrs-core/src/cuda/hc_fused.rs @@ -444,6 +444,52 @@ mod tests { } } + /// Source-level tripwires, mirroring `sinkhorn::tests`: the IEEE + /// intrinsics and the fast-math `#error` guard must stay in hc_fused.cu, + /// fast-math approximations must stay out, and build.rs must keep the file + /// out of the `--use_fast_math` glob and in the `--fmad=false` builder. + /// Getting this wiring wrong is silent — the kernel still runs, it just + /// stops being bit-identical — so the `#error` guard is the hard stop and + /// these string checks catch it on CPU CI too. + #[test] + fn kernel_source_and_build_wiring_guards() { + let cu = include_str!("hc_fused.cu"); + assert!( + cu.contains("#if defined(__USE_FAST_MATH__)") && cu.contains("#error"), + "hc_fused.cu lost its fast-math #error guard" + ); + for required in ["__fadd_rn", "__fmul_rn", "candle_recip", "candle_sigmoid"] { + assert!(cu.contains(required), "hc_fused.cu lost required token {required}"); + } + for forbidden in ["__expf(", "__logf(", "__fdividef(", "rsqrtf(", "__frsqrt_rn("] { + assert!( + !cu.contains(forbidden), + "hc_fused.cu contains {forbidden}, which is not what candle-kernels computes" + ); + } + + let build = include_str!("../../build.rs"); + let exclude = build + .split(".exclude(&[") + .nth(1) + .and_then(|s| s.split("])").next()) + .expect("build.rs no longer calls .exclude(&[..]) on the fast-math builder"); + assert!( + exclude.contains("\"hc_fused.cu\""), + "build.rs no longer excludes hc_fused.cu from the fast-math builder \ + (exclude list is: {exclude})" + ); + assert!( + build.contains(r#""src/cuda/hc_fused.cu""#), + "build.rs no longer feeds hc_fused.cu to the IEEE (no-fast-math) builder" + ); + assert!( + build.contains("--fmad=false"), + "build.rs lost --fmad=false, so nvcc may contract mul+add into an FMA and break \ + bit-identity with candle's unfused op chain" + ); + } + #[test] fn candle_reduce_block_dim_matches_fast_reduce() { // cuda_backend/mod.rs: usize::min(1024, el_to_sum_per_block).next_power_of_two() diff --git a/mistralrs-core/src/cuda/sinkhorn.rs b/mistralrs-core/src/cuda/sinkhorn.rs index 40173853b..c101c7d4c 100644 --- a/mistralrs-core/src/cuda/sinkhorn.rs +++ b/mistralrs-core/src/cuda/sinkhorn.rs @@ -474,9 +474,23 @@ mod tests { } let build = include_str!("../../build.rs"); + // The exclude list grows as more bit-identity-critical kernels are + // added (hc_fused.cu joined it), so match on the list's CONTENTS rather + // than on its exact spelling — otherwise this guard fails for the one + // reason that is not a regression, and gets "fixed" by weakening it. + let exclude = build + .split(".exclude(&[") + .nth(1) + .and_then(|s| s.split("])").next()) + .expect("build.rs no longer calls .exclude(&[..]) on the fast-math builder"); assert!( - build.contains(r#".exclude(&["sinkhorn.cu"])"#), - "build.rs no longer excludes sinkhorn.cu from the fast-math builder" + exclude.contains("\"sinkhorn.cu\""), + "build.rs no longer excludes sinkhorn.cu from the fast-math builder \ + (exclude list is: {exclude})" + ); + assert!( + build.contains(r#""src/cuda/sinkhorn.cu""#), + "build.rs no longer feeds sinkhorn.cu to the IEEE (no-fast-math) builder" ); assert!( build.contains("--fmad=false"), From ab5fbd54d38b228fe0438b2b2ddbae95567016a3 Mon Sep 17 00:00:00 2001 From: Nirupam Bhowmick Date: Wed, 19 Aug 2026 22:32:41 +0100 Subject: [PATCH 22/22] docs(ArcMoE): cite the measured H200 numbers, not the estimates Measured on H200, git 6f4cd0dd, bin f6ec4e85, toggled with ARC_HC_FUSED on one binary so the comparison has no second variable: ARC_HC_FUSED=0 7,494.3 launches/step 10,892.4 allocs/step 17.3 tok/s 57.70 ms/T ARC_HC_FUSED=1 5,691.0 launches/step 8,918.5 allocs/step 20.1 tok/s 49.66 ms/T The off leg reproduces origin/master (7,477.9 launches, 17.0 tok/s) to within 0.2%, so the switch is the only thing that moved and nothing else regressed. Router region 76 -> 38 kernels/layer, 0.207 -> 0.100 ms GPU. Sinkhorn 30.03 -> 9.12 us/call (2.583 -> 0.785 ms/step). Output bit-identical across the toggle. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_01SpVNMpb13HkUXqSqbN1o9H --- mistralrs-core/src/cuda/hc_fused.cu | 22 +++++++++++++++------- mistralrs-core/src/models/dsv4_mhc.rs | 8 +++++--- 2 files changed, 20 insertions(+), 10 deletions(-) diff --git a/mistralrs-core/src/cuda/hc_fused.cu b/mistralrs-core/src/cuda/hc_fused.cu index bdfc711ce..a189822aa 100644 --- a/mistralrs-core/src/cuda/hc_fused.cu +++ b/mistralrs-core/src/cuda/hc_fused.cu @@ -5,13 +5,21 @@ // --------------------------------------------------------------------------- // WHY // --------------------------------------------------------------------------- -// At b=1 the V4 decode step is launch-bound, not bandwidth-bound: the measured -// step issues ~7.9k `cuLaunchKernel` calls, of which the router region alone is -// 3,259 (41.8%) spread over 86 contiguous spans -- exactly two per layer -// (`mhc_attn_pre` and `mhc_ffn_pre`, 43 layers). Those spans operate on -// `[1, 24]` / `[1, 4, 4]` / `[1, 256]` tensors: mean kernel duration 2.79 us, -// 89% of them under 5 us. They are pure launch overhead wearing a kernel -// costume. +// At b=1 the V4 decode step is launch-bound, not bandwidth-bound. Measured on +// an H200 at master ef581e9c8: 7,477.9 `cuLaunchKernel` calls per decode step +// costing 18.976 ms/step of host time, of which the router region is 76 kernels +// per layer (3,268/step, 43.7%) spread over 86 contiguous spans -- exactly two +// per layer (`mhc_attn_pre` and `mhc_ffn_pre`, 43 layers). Those spans operate +// on `[1, 24]` / `[1, 4, 4]` / `[1, 256]` tensors: mean kernel duration 2.72 us, +// 89% of them under 5 us, 28.4% of all kernel time. They are pure launch +// overhead wearing a kernel costume. +// +// MEASURED RESULT of this file plus the sinkhorn fix, same box, same binary, +// toggled with ARC_HC_FUSED: 7,494.3 -> 5,691.0 launches/step (-24.1%), +// 10,892.4 -> 8,918.5 allocations/step, 57.70 -> 49.66 ms/token +// (17.3 -> 20.1 tok/s). The router region itself goes 76 -> 38 kernels/layer +// and 0.207 -> 0.100 ms of GPU time. Generated tokens and logprobs are +// bit-identical across the toggle (6 prompts, 768 logprob values, 0 mismatches). // // Two expressions dominate the count and are collapsed here: // diff --git a/mistralrs-core/src/models/dsv4_mhc.rs b/mistralrs-core/src/models/dsv4_mhc.rs index d3c740ed7..57fb1e2d7 100644 --- a/mistralrs-core/src/models/dsv4_mhc.rs +++ b/mistralrs-core/src/models/dsv4_mhc.rs @@ -290,10 +290,12 @@ impl V4MHCLayerParams { // statistic (`sqr -> fast_sum -> affine -> affine -> recip -> sqrt -> // bmul`) plus 11 for the three scoring blocks — all on 24 floats once // the reduction is done, twice per layer, 43 layers. At b=1 that is - // ~1,460 of the step's ~7,900 kernel launches, i.e. pure launch - // overhead. `cuda/hc_fused.cu` is bit-identical to this chain by + // ~1,460 of the decode step's measured 7,494 kernel launches, i.e. pure + // launch overhead. `cuda/hc_fused.cu` is bit-identical to this chain by // construction, not by tolerance; the eager path stays reachable via - // `ARC_HC_FUSED=0` so the two can be A/B'd from one binary. + // `ARC_HC_FUSED=0` so the two can be A/B'd from one binary — and they + // were: flipping it moves 1,803 launches/step and 8.04 ms/token while + // leaving 6 greedy completions and their 768 logprobs bit-identical. let on_cuda = crate::cuda::hc_fused::usable(&x_flat); let shapes_ok = x_flat.is_contiguous() && mixes_raw.is_contiguous()