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", ] diff --git a/Cargo.toml b/Cargo.toml index b8dde1f07..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 = "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 = "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" 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; +} 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..39bb43314 --- /dev/null +++ b/arc-tools/kv_fp8_count_per_step.py @@ -0,0 +1,125 @@ +#!/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]}") + + # ---- 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() diff --git a/arc-tools/kv_fp8_nsys_ab.sh b/arc-tools/kv_fp8_nsys_ab.sh new file mode 100755 index 000000000..80a38692e --- /dev/null +++ b/arc-tools/kv_fp8_nsys_ab.sh @@ -0,0 +1,81 @@ +#!/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). +# +# 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:-} +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 +} + +exec 9>"$LOCK" +for arm in before:cpu after:fused; do + A=${arm%%:*} + M=${arm##*:} + 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_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 \ + --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=$? + 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; } + [ -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 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() 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/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..a189822aa --- /dev/null +++ b/mistralrs-core/src/cuda/hc_fused.cu @@ -0,0 +1,299 @@ +// 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. 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: +// +// 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..7376ce005 --- /dev/null +++ b/mistralrs-core/src/cuda/hc_fused.rs @@ -0,0 +1,638 @@ +//! 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"), + }; + // 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 + /// `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 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(( + 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))), + )) + } + + /// 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) + } + } + + /// 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() + 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/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"), 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_kv_fp8.rs b/mistralrs-core/src/models/dsv4_kv_fp8.rs index 5752bebf5..a52feac5d 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,33 +62,145 @@ 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_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 /// 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, } +/// 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(|| { - if std::env::var_os("ARC_GPU_ACT_QUANT").is_some() { - Self::GpuApprox - } else { - Self::CpuExact + 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 }) } } +/// 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 +278,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 +370,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 +399,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 +659,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] @@ -506,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" + ); + } + } } diff --git a/mistralrs-core/src/models/dsv4_mhc.rs b/mistralrs-core/src/models/dsv4_mhc.rs index 00afc5940..57fb1e2d7 100644 --- a/mistralrs-core/src/models/dsv4_mhc.rs +++ b/mistralrs-core/src/models/dsv4_mhc.rs @@ -278,57 +278,111 @@ 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 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 — 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() + && 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] diff --git a/mistralrs-core/src/pipeline/normal.rs b/mistralrs-core/src/pipeline/normal.rs index d412a7cd9..0bcc24e8a 100644 --- a/mistralrs-core/src/pipeline/normal.rs +++ b/mistralrs-core/src/pipeline/normal.rs @@ -98,6 +98,175 @@ 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") +} + +/// 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 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 { + 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, @@ -1702,15 +1871,37 @@ 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); + 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: // drives RoPE + the fixed-capacity KV write slot, and // makes warmup forwards take the shape-constant path so @@ -2528,3 +2719,132 @@ 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 + ); + } +} + +/// 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_unparsable_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" + ); + } +} 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 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..2f470e95e --- /dev/null +++ b/mistralrs-quant/kernels/arc_kvquant/arc_kvquant.cu @@ -0,0 +1,496 @@ +// 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`. 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` +// 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`). +// --------------------------------------------------------------------------- +// --------------------------------------------------------------------------- +// 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); +} +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 `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 arc_fadd(arc_fmul(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(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 + // 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(arc_fmul(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(arc_fdiv(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;