Skip to main content

cortiq_engine/
gpu.rs

1//! Facade for GPU backends: a single call entry point for qtensor/pipeline/
2//! linear_core. Job types and the threshold are canonical HERE; behind the
3//! facade dispatch goes to a platform backend:
4//!   - `gpu_metal` (Apple Silicon, unified memory + no-copy buffers);
5//!   - `gpu_wgpu` (C1: Vulkan/DX12/Metal — NVIDIA/Radeon/Intel/Apple,
6//!     weights resident in VRAM), available under `--features gpu`.
7//!
8//! Runtime selection via `CMF_GPU`: `1` — native Metal (macOS) or wgpu
9//! (other OSes); `wgpu` — force wgpu (including for the local
10//! Metal-via-wgpu parity test). Any backend refusal — `false` and the honest
11//! CPU path, no partial results.
12
13use cortiq_core::CmfModel;
14use std::cell::Cell;
15use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU32, AtomicU64, Ordering};
16use std::sync::{Arc, OnceLock};
17
18thread_local! {
19    /// Index of the current forward layer (−1 = outside a numbered layer:
20    /// lm_head/embed — always allowed). The pipeline sets it before
21    /// each layer so that the GPU/CPU layer-split works.
22    static CUR_LAYER: Cell<i64> = const { Cell::new(-1) };
23    /// Inside `cpu_scope` every GPU gate reports disabled: the timed CPU
24    /// arm of a probe (and a class that lost its probe) must run PURE
25    /// CPU, or inner per-op hooks would re-enter the GPU and poison the
26    /// comparison.
27    static CPU_ONLY: Cell<bool> = const { Cell::new(false) };
28    /// "This op paid a one-off cost" (weight upload / first pipeline
29    /// build): backends set it, `probe_record` discards the sample so
30    /// only steady-state timings compete.
31    static PROBE_COLD: Cell<bool> = const { Cell::new(false) };
32}
33
34/// Run `f` with the GPU gates off on this thread (pure-CPU arm).
35pub fn cpu_scope<R>(f: impl FnOnce() -> R) -> R {
36    struct Restore(bool);
37    impl Drop for Restore {
38        fn drop(&mut self) {
39            CPU_ONLY.with(|c| c.set(self.0));
40        }
41    }
42    let previous = CPU_ONLY.with(|c| c.replace(true));
43    let _restore = Restore(previous);
44    f()
45}
46
47/// Backends: name the device once at init. The probe cache is keyed by
48/// it, because a verdict is a property of THIS silicon and nothing else.
49/// First writer wins: a process runs one backend, and on the rare host
50/// where two initialize, the one that came up first is the one in use.
51pub fn probe_set_device(label: &str) {
52    let _ = DEVICE_LABEL.set(label.to_string());
53}
54
55fn device_label() -> &'static str {
56    DEVICE_LABEL.get().map(String::as_str).unwrap_or("unknown")
57}
58
59static DEVICE_LABEL: std::sync::OnceLock<String> = std::sync::OnceLock::new();
60
61/// Somewhere this process may write small caches.
62///
63/// `std::env::temp_dir()` is NOT that place on Android: with no `TMPDIR`
64/// it answers `/tmp`, which does not exist in an app sandbox, and every
65/// write fails silently — measured, after the pipeline cache appeared to
66/// work in a shell (where `TMPDIR=/data/local/tmp`) and did nothing at
67/// all in the app. The loader points this at the model's own directory,
68/// which is somewhere the caller already writes.
69static CACHE_DIR: std::sync::OnceLock<std::path::PathBuf> = std::sync::OnceLock::new();
70
71/// Loader: name a directory this process can write to. First call wins.
72pub fn set_cache_dir(dir: std::path::PathBuf) {
73    let _ = CACHE_DIR.set(dir);
74}
75
76/// Same directory, for the backends.
77pub fn cache_dir_pub() -> std::path::PathBuf {
78    cache_dir()
79}
80
81fn cache_dir() -> std::path::PathBuf {
82    if let Some(d) = CACHE_DIR.get() {
83        return d.clone();
84    }
85    match std::env::var_os("TMPDIR") {
86        Some(t) => std::path::PathBuf::from(t),
87        None => std::env::temp_dir(),
88    }
89}
90
91/// Where decided verdicts are remembered between runs. `CMF_PROBE_CACHE`
92/// overrides the path; `0` disables the cache entirely.
93fn probe_cache_path() -> Option<std::path::PathBuf> {
94    match std::env::var("CMF_PROBE_CACHE") {
95        Ok(v) if v == "0" => None,
96        Ok(v) => Some(std::path::PathBuf::from(v)),
97        Err(_) => Some(cache_dir().join("cortiq-gpu-probe.tsv")),
98    }
99}
100
101/// One line per decided class: `version \t device \t class \t winner`.
102/// A different engine build or a different device simply does not match,
103/// so a stale file is inert rather than wrong.
104fn probe_cache_key_named(class: &str) -> String {
105    format!(
106        "{}\t{}\t{}",
107        env!("CARGO_PKG_VERSION"),
108        device_label(),
109        class
110    )
111}
112
113const CLASS_NAMES: [&str; 7] = [
114    "ffn",
115    "matvec",
116    "matmat",
117    "qkv-batch",
118    "matmat-wide",
119    "lm-head",
120    "gemm-nt",
121];
122
123/// Adopt every verdict this device already reached in an earlier run.
124///
125/// Probing is not cheap and it is not free of consequences: on a
126/// Snapdragon 778G the three deciding classes took **three minutes of
127/// wall clock** before the first token, every process, and in the phone
128/// app that was the whole first answer — 209.6 s for 25 tokens against
129/// 10.5 s on the CPU path. The verdict itself was the same every time.
130/// Paying to rediscover it is the defect; the answer is to write it down.
131fn probe_cache_load() {
132    static ONCE: std::sync::Once = std::sync::Once::new();
133    ONCE.call_once(|| {
134        let Some(path) = probe_cache_path() else {
135            return;
136        };
137        // Unit tests share this process and its default cache path; a
138        // verdict left by an earlier run would decide a class before the
139        // arbitration tests get to watch it alternate. Tests that mean to
140        // exercise the cache point `CMF_PROBE_CACHE` at their own file.
141        if cfg!(test) && std::env::var("CMF_PROBE_CACHE").is_err() {
142            return;
143        }
144        let Ok(text) = std::fs::read_to_string(&path) else {
145            return;
146        };
147        probe_cache_adopt(&text);
148    });
149}
150
151/// Apply verdicts from a cache file's text. Split out from the file
152/// reading so the adoption rule — including which lines must be IGNORED
153/// — is testable without a filesystem.
154fn probe_cache_adopt(text: &str) {
155    for line in text.lines() {
156        let Some((key, verdict)) = line.rsplit_once('\t') else {
157            continue;
158        };
159        let winner = match verdict.trim() {
160            "gpu" => 1u8,
161            "cpu" => 2u8,
162            _ => continue,
163        };
164        for (i, name) in CLASS_NAMES.iter().enumerate() {
165            if probe_cache_key_named(name) == key {
166                let _ = PROBES[i].state.compare_exchange(
167                    0,
168                    winner,
169                    Ordering::Relaxed,
170                    Ordering::Relaxed,
171                );
172                tracing::debug!("gpu probe [{name}]: remembered → {verdict}");
173            }
174        }
175    }
176}
177
178/// Remember a verdict for the next run. Best-effort: a read-only cache
179/// directory costs a re-probe, never a failure.
180fn probe_cache_store(c: OpClass, winner: u8) {
181    let Some(path) = probe_cache_path() else {
182        return;
183    };
184    let line = format!(
185        "{}\t{}\n",
186        probe_cache_key_named(CLASS_NAMES[c as usize]),
187        if winner == 1 { "gpu" } else { "cpu" }
188    );
189    use std::io::Write;
190    if let Ok(mut f) = std::fs::OpenOptions::new()
191        .create(true)
192        .append(true)
193        .open(&path)
194    {
195        let _ = f.write_all(line.as_bytes());
196    }
197}
198
199/// Backends: note a one-off cost (weight upload, buffer-cache fill) so
200/// the probe discards this sample.
201/// Every buffer creation anywhere bumps this; the graph's bind-group
202/// cache treats any cold event as total invalidation — a stale bind
203/// group is silent corruption, a cleared cache is one re-encoded token.
204pub fn cold_epoch() -> u64 {
205    COLD_EPOCH.load(std::sync::atomic::Ordering::Relaxed)
206}
207static COLD_EPOCH: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
208
209pub(crate) fn probe_note_cold() {
210    COLD_EPOCH.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
211    PROBE_COLD.with(|c| c.set(true));
212}
213
214/// Peek the cold flag without consuming it (`probe_record` consumes).
215/// Contention heuristics use this: a slow COLD op is a one-off build
216/// cost, not evidence the device is busy.
217pub(crate) fn probe_was_cold() -> bool {
218    PROBE_COLD.with(|c| c.get())
219}
220
221/// Pipeline: mark the current layer (or −1 outside layers) for layer-split.
222pub fn set_layer(l: i64) {
223    CUR_LAYER.with(|c| c.set(l));
224}
225
226/// The layer `set_layer` last marked on this thread (−1 outside layers).
227pub fn cur_layer() -> i64 {
228    CUR_LAYER.with(|c| c.get())
229}
230
231/// Parse `CMF_GPU_LAYERS` («0-19», «0,2,4», «0-9,30-39») once.
232/// None = no restriction (all layers on GPU). Garbage → also no restriction.
233fn layer_ranges() -> &'static Option<Vec<(i64, i64)>> {
234    static R: OnceLock<Option<Vec<(i64, i64)>>> = OnceLock::new();
235    R.get_or_init(|| {
236        let s = std::env::var("CMF_GPU_LAYERS").ok()?;
237        let mut v = Vec::new();
238        for part in s.split(',') {
239            let part = part.trim();
240            match part.split_once('-') {
241                Some((a, b)) => v.push((a.trim().parse().ok()?, b.trim().parse().ok()?)),
242                None => {
243                    let x: i64 = part.parse().ok()?;
244                    v.push((x, x));
245                }
246            }
247        }
248        Some(v)
249    })
250}
251
252fn layer_allowed() -> bool {
253    match layer_ranges() {
254        None => true,
255        Some(ranges) => {
256            let cur = CUR_LAYER.with(|c| c.get());
257            cur < 0 || ranges.iter().any(|(a, b)| cur >= *a && cur <= *b)
258        }
259    }
260}
261
262/// GPU allowed FOR THE CURRENT LAYER: backend is initialized AND the layer
263/// falls within `CMF_GPU_LAYERS` (GPU/CPU layer-split) AND we are not
264/// inside a `cpu_scope`. Op gates call this.
265pub fn enabled_here() -> bool {
266    !CPU_ONLY.with(|c| c.get()) && enabled() && layer_allowed()
267}
268
269// ── Runtime GPU-vs-CPU probe ────────────────────────────────────────────
270// CMF_GPU=1 does not TRUST that the device wins — it MEASURES. For each
271// op class the first calls alternate arms: GPU timed vs pure-CPU timed
272// (under cpu_scope). Cold GPU calls (weight upload / cache fill) are
273// discarded; after PROBE_SAMPLES clean samples per arm the faster arm is
274// chosen for the rest of the process. Rationale: submit+poll latency
275// differs by an order of magnitude across driver stacks (Metal/PCIe
276// ~3-4 ms, Vulkan/4090 ~0.3 ms) — a static threshold cannot know whether
277// per-op offload pays off HERE. CMF_GPU_PROBE=0 → always trust the GPU.
278
279/// GPU-eligible op classes, each with an independent probe.
280#[derive(Clone, Copy)]
281pub enum OpClass {
282    /// Whole FFN chain in one submission (dense / MoE block).
283    Ffn = 0,
284    /// Large hybrid CPU∥GPU matvec (lm_head class).
285    Matvec = 1,
286    /// Prefill GEMM (matmat).
287    Matmat = 2,
288    /// Batched matvecs of one input (QKV).
289    Batch = 3,
290    /// Prefill GEMM at image-diffusion widths (b ≥ 128). Probed apart
291    /// from `Matmat`: one imagegen process runs BOTH populations
292    /// (prompt encode b≈40 where the GPU wins big, DiT b≥256 where
293    /// the CPU AMX arm is competitive) — a single shared verdict locks
294    /// the wrong arm for whichever population samples second.
295    MatmatWide = 4,
296    /// The lm_head itself, apart from the merely-large matvecs. Same
297    /// reasoning as `MatmatWide`, and DeepSeek-V4 is where it bit: its
298    /// attention projections are 37M weights and its head is 529M, so
299    /// the projections' verdict — CPU, honestly measured at 0.19 ms —
300    /// decided for a matvec fourteen times their size that took 11 ms
301    /// a token on the host.
302    MatvecHead = 5,
303    /// The blocked f32 GEMM (`fcd_ops::gemm_nt`): attention's QKᵀ and
304    /// AV, and the VAE decoders' projections. It used to take every job
305    /// over 4 M MACs on sight, with no CPU arm to lose to — which on
306    /// the MiniMax-H3 video decoder was three times SLOWER than the
307    /// host it displaced. Its population is per-head slices, nothing
308    /// like the weight GEMMs above, so it probes on its own.
309    GemmNt = 6,
310}
311
312/// Which probe a large matvec belongs to. The head is an order of
313/// magnitude bigger than anything else that reaches this gate, and the
314/// two populations do not have the same answer.
315pub fn matvec_class(rows: usize, cols: usize) -> OpClass {
316    if rows * cols >= 67_108_864 {
317        OpClass::MatvecHead
318    } else {
319        OpClass::Matvec
320    }
321}
322
323/// Probe verdict for one call.
324pub enum ProbeArm {
325    /// Run the GPU path (during probing: timed, recorded).
326    Gpu,
327    /// Probing: run the CPU path under `cpu_scope`, timed, recorded.
328    CpuTimed,
329    /// Decided: CPU won — run the CPU path (under `cpu_scope`).
330    Cpu,
331}
332
333/// Clean samples per arm before a class decides.
334const PROBE_SAMPLES: u32 = 6;
335
336/// Declines before a class gives the work to the host for good. High
337/// enough that a transient refusal — an unsealed state during prefill, a
338/// shape the kernel skips this once — cannot settle the question.
339const PROBE_DECLINE_LIMIT: u32 = 16;
340
341/// Device samples discarded before any count — see `Probe::gpu_burn`.
342const PROBE_WARMUP: u32 = 1;
343
344struct Probe {
345    /// 0 = probing, 1 = GPU won, 2 = CPU won.
346    state: AtomicU8,
347    flip: AtomicU32,
348    gpu_ns: AtomicU64,
349    gpu_n: AtomicU32,
350    /// Times the device arm was chosen and the device DECLINED.
351    ///
352    /// A decline carries no timing, so nothing is recorded — and a class
353    /// whose device path always refuses therefore never reaches a
354    /// verdict, alternates arms forever, and pays a failed device
355    /// attempt on half of every token's calls. Measured on an M4 with
356    /// LFM2.5-2.6B: `ffn` was still undecided after 9000 calls, and a
357    /// token cost 83.55 ms against 41.85 with the device off — twice the
358    /// price for work the host did anyway.
359    declines: AtomicU32,
360    /// GPU samples still to discard as warm-up.
361    ///
362    /// The cold flag catches buffer and weight uploads, but a compute
363    /// pipeline is compiled on first use and not every creation site
364    /// raises it — the wgpu path has 21 pipeline creations against 12
365    /// cold notes. One uncaught shader compile is enough to lose a
366    /// class for the whole process: `gemm-nt` on an A100 was recorded at
367    /// 117.01 ms against the host's 3.19 and sent to the CPU, which
368    /// parked a 27B bake on 2.6 cores with the card idle. The decision
369    /// already uses each arm's BEST sample, so discarding the first
370    /// GPU sample costs one extra round trip and removes the whole
371    /// class of first-call artefacts.
372    gpu_burn: AtomicU32,
373    cpu_ns: AtomicU64,
374    cpu_n: AtomicU32,
375    /// Best (minimum) sample per arm. The DECISION compares these:
376    /// means are poisoned by one-off cold costs the cold-flag cannot
377    /// see — e.g. the CPU arm's first mmap-cold expert matvec page
378    /// faults its weights in and reads 3× its steady state, which
379    /// locked the GPU arm on a 35B MoE at a 4× real-world loss. The
380    /// minimum is each arm's honest steady-state pace.
381    gpu_min: AtomicU64,
382    cpu_min: AtomicU64,
383}
384
385impl Probe {
386    const fn new() -> Self {
387        Self {
388            state: AtomicU8::new(0),
389            flip: AtomicU32::new(0),
390            gpu_ns: AtomicU64::new(0),
391            gpu_n: AtomicU32::new(0),
392            declines: AtomicU32::new(0),
393            gpu_burn: AtomicU32::new(PROBE_WARMUP),
394            cpu_ns: AtomicU64::new(0),
395            cpu_n: AtomicU32::new(0),
396            gpu_min: AtomicU64::new(u64::MAX),
397            cpu_min: AtomicU64::new(u64::MAX),
398        }
399    }
400}
401
402static PROBES: [Probe; 7] = [
403    Probe::new(),
404    Probe::new(),
405    Probe::new(),
406    Probe::new(),
407    Probe::new(),
408    Probe::new(),
409    Probe::new(),
410];
411
412/// A caller that knows its loop is long, uniform and warm can say so: the
413/// probe times ops in isolation and alternates arms to do it, which reads a
414/// sustained diffusion step as slower on the device than it is. Measured on
415/// an M4 at 672 video tokens: the probe picked the CPU at 1.25 ms against
416/// 0.88 ms per op, and the loop it picked for ran 23.9 s a step against the
417/// device's 19.7 s.
418static TRUST_GPU: AtomicBool = AtomicBool::new(false);
419
420/// Take the probe out of the loop until the guard drops.
421pub fn trust_gpu() -> GpuTrust {
422    let was = TRUST_GPU.swap(true, Ordering::Relaxed);
423    GpuTrust(was)
424}
425
426pub struct GpuTrust(bool);
427
428impl Drop for GpuTrust {
429    fn drop(&mut self) {
430        TRUST_GPU.store(self.0, Ordering::Relaxed);
431    }
432}
433
434fn probe_on_for(c: OpClass) -> bool {
435    // The trust is only for the *wide* class. A sustained diffusion step is
436    // where the probe reads a warm device as cold; the narrow batches inside
437    // the same loop — an audio stream of fifty-one tokens against the same
438    // weights — are small enough that submit latency can genuinely beat the
439    // arithmetic, and there the probe is right and should keep deciding.
440    if TRUST_GPU.load(Ordering::Relaxed) && matches!(c, OpClass::MatmatWide | OpClass::Ffn) {
441        return false;
442    }
443    probe_on()
444}
445
446fn probe_on() -> bool {
447    static ON: OnceLock<bool> = OnceLock::new();
448    *ON.get_or_init(|| {
449        std::env::var("CMF_GPU_PROBE")
450            .map(|v| v != "0" && v != "off")
451            .unwrap_or(true)
452    })
453}
454
455/// q1 ops on the native Metal backend skip the probe entirely: the CPU
456/// q1 kernel is load-port-bound, the GPU one wins warm — and probe
457/// alternation itself cools the device between samples (measured: block
458/// times 5.8 ms warm vs 8.8 ms mixed). Other backends keep probing.
459pub fn q1_force() -> bool {
460    #[cfg(target_os = "macos")]
461    {
462        backend() == Backend::Metal
463    }
464    #[cfg(not(target_os = "macos"))]
465    {
466        false
467    }
468}
469
470/// Should a FUSED whole-block path trust the device instead of asking
471/// the per-op probe? True on native Metal and on discrete wgpu adapters.
472///
473/// The probe answers "is one wide matmat faster on the GPU", and for the
474/// DiT on Metal that is a coin flip — measured 2.62 ms GPU vs 2.56 ms
475/// CPU, a 2% spread that lands on either arm run to run. But the fused
476/// block's advantage is not per-op speed, it is that the hidden state,
477/// the packs and the attention panels never leave the device: end to end
478/// the whole-block path renders a 512² Lumina step in ~5.4 s against
479/// ~8.4 s when the probe happens to pick the CPU. Gating a fusion win on
480/// a per-op tie made every second render half-speed at random.
481///
482/// On a discrete card the verdict is never in doubt — an RTX 3090 against
483/// a 256-core EPYC measured 11.5 ms vs 31 ms per wide op, four runs out
484/// of four — so the probe's sampling phase is pure cost: it alone was 10%
485/// of a 512² render (74.3 s against 66.9 s with the probe off). Integrated
486/// and mobile adapters keep probing; there the submit latency is real and
487/// can genuinely lose.
488pub fn fused_block_trusted() -> bool {
489    #[cfg(target_os = "macos")]
490    if backend() == Backend::Metal {
491        return true;
492    }
493    wgpu_graph_default()
494}
495
496/// Which arm should this GPU-eligible call take? Consult AFTER the
497/// eligibility gates (`enabled_here` / `min_rows`) so only real
498/// candidates alternate.
499/// While a class is still probing, a call whose weights are NOT yet on
500/// the card should take the GPU arm anyway: the upload is work the next
501/// step needs regardless, and the sample it produces is discarded as
502/// cold — so handing that call to the CPU arm buys nothing and costs a
503/// host GEMM. Measured on a diffusion stack, where every layer is
504/// touched once per step and therefore EVERY first-step GPU sample is
505/// cold: one projection drew the CPU arm for the whole first step, 9.8 s
506/// against the 2.8 s it costs once the weights are warm.
507pub fn weight_is_resident(model: &Arc<CmfModel>, idx: usize) -> bool {
508    #[cfg(all(feature = "gpu", not(target_os = "macos")))]
509    {
510        return crate::gpu_wgpu::weight_is_resident(model, idx);
511    }
512    #[cfg(not(all(feature = "gpu", not(target_os = "macos"))))]
513    {
514        let _ = (model, idx);
515        true
516    }
517}
518
519pub fn probe_arm_cold_prefers_gpu(c: OpClass, weights_resident: bool) -> ProbeArm {
520    if !weights_resident && probe_deciding(c) {
521        return ProbeArm::Gpu;
522    }
523    probe_arm(c)
524}
525
526pub fn probe_arm(c: OpClass) -> ProbeArm {
527    // Every arbitrated call starts with a clean cold flag: both the
528    // sample discard in `probe_record` and the contention kill-switch
529    // read it AFTER the op, so a stale note from a previous call on
530    // this thread must not leak in.
531    PROBE_COLD.with(|f| f.set(false));
532    if !probe_on_for(c) {
533        return ProbeArm::Gpu;
534    }
535    probe_cache_load();
536    let p = &PROBES[c as usize];
537    match p.state.load(Ordering::Relaxed) {
538        1 => ProbeArm::Gpu,
539        2 => ProbeArm::Cpu,
540        _ => {
541            if p.flip.fetch_add(1, Ordering::Relaxed) % 2 == 0 {
542                ProbeArm::Gpu
543            } else {
544                ProbeArm::CpuTimed
545            }
546        }
547    }
548}
549
550/// The device arm was chosen and the device refused the work, so there
551/// is no time to record. Callers that fall through to the host MUST say
552/// so here, or the class can never decide.
553pub fn probe_note_decline(c: OpClass) {
554    let p = &PROBES[c as usize];
555    if p.state.load(Ordering::Relaxed) != 0 {
556        return;
557    }
558    let n = p.declines.fetch_add(1, Ordering::Relaxed) + 1;
559    if n >= PROBE_DECLINE_LIMIT
560        && p.state
561            .compare_exchange(0, 2, Ordering::Relaxed, Ordering::Relaxed)
562            .is_ok()
563    {
564        tracing::info!(
565            "gpu probe [{}]: device declined {n} times → cpu",
566            CLASS_NAMES[c as usize]
567        );
568    }
569}
570
571/// Record a timed arm sample; on the `PROBE_SAMPLES`-th clean sample of
572/// BOTH arms the class decides for the rest of the process.
573pub fn probe_record(c: OpClass, gpu: bool, dur: std::time::Duration) {
574    probe_record_into(
575        &PROBES[c as usize],
576        CLASS_NAMES[c as usize],
577        Some(c),
578        gpu,
579        dur,
580    )
581}
582
583/// The body of `probe_record` over ONE probe, so the decision can be
584/// driven in a test without touching the process-wide array.
585fn probe_record_into(
586    p: &Probe,
587    class_name: &str,
588    cache: Option<OpClass>,
589    gpu: bool,
590    dur: std::time::Duration,
591) {
592    if p.state.load(Ordering::Relaxed) != 0 {
593        return;
594    }
595    if gpu && PROBE_COLD.with(|f| f.replace(false)) {
596        return; // one-off cost in this call — not a steady-state sample
597    }
598    if gpu {
599        // Load-then-store rather than fetch_sub: a blind decrement at
600        // zero wraps a u32 to its maximum and mutes the arm forever.
601        // A benign race here burns one extra sample, which is free.
602        let left = p.gpu_burn.load(Ordering::Relaxed);
603        if left > 0 {
604            p.gpu_burn.store(left - 1, Ordering::Relaxed);
605            return; // warm-up: the first device sample builds its pipeline
606        }
607    }
608    let ns = dur.as_nanos().min(u64::MAX as u128) as u64;
609    if gpu {
610        p.gpu_ns.fetch_add(ns, Ordering::Relaxed);
611        p.gpu_n.fetch_add(1, Ordering::Relaxed);
612        p.gpu_min.fetch_min(ns, Ordering::Relaxed);
613    } else {
614        p.cpu_ns.fetch_add(ns, Ordering::Relaxed);
615        p.cpu_n.fetch_add(1, Ordering::Relaxed);
616        p.cpu_min.fetch_min(ns, Ordering::Relaxed);
617    }
618    let (gn, cn) = (
619        p.gpu_n.load(Ordering::Relaxed),
620        p.cpu_n.load(Ordering::Relaxed),
621    );
622    if gn >= 2 && cn >= 2 {
623        // Decide on each arm's BEST sample — the steady-state pace.
624        // Means carry one-off cold costs (mmap page-in on the CPU arm)
625        // that the cold-flag machinery cannot see.
626        let g = p.gpu_min.load(Ordering::Relaxed) as f64;
627        let cp = p.cpu_min.load(Ordering::Relaxed) as f64;
628        // Early verdict on a ≥2× gap — no reason to keep feeding the
629        // losing arm; close races take the full sample count. It was 3×,
630        // and the cost of that half-octave was measured: a DiT whose
631        // wide GEMMs run 11.4 ms on the device against 32.2 on the host
632        // (2.8×) kept ALTERNATING through the whole diffusion stack, and
633        // because the alternation counter is shared per class in call
634        // order, one projection drew the CPU arm every single time — 9.9
635        // seconds a step on a kernel that needs 0.4. Both arms are
636        // compared on their BEST sample, so a 2× gap is not noise.
637        if (gn < PROBE_SAMPLES || cn < PROBE_SAMPLES) && g < cp * 2.0 && cp < g * 2.0 {
638            return;
639        }
640        let winner = if g <= cp { 1 } else { 2 };
641        if p.state
642            .compare_exchange(0, winner, Ordering::Relaxed, Ordering::Relaxed)
643            .is_ok()
644        {
645            tracing::info!(
646                "gpu probe [{}]: gpu {:.2} ms vs cpu {:.2} ms per op → {}",
647                class_name,
648                g / 1e6,
649                cp / 1e6,
650                if winner == 1 { "gpu" } else { "cpu" },
651            );
652            if let Some(c) = cache {
653                probe_cache_store(c, winner);
654            }
655        }
656    }
657}
658
659/// Is the class still collecting samples? (Call sites use this to route
660/// cold-weight calls away from the GPU arm during probing.)
661pub fn probe_deciding(c: OpClass) -> bool {
662    probe_on_for(c) && PROBES[c as usize].state.load(Ordering::Relaxed) == 0
663}
664
665/// Probing helper: true — tensor `idx`'s quant weights are ALREADY
666/// device-resident (a clean GPU sample is possible now); false — they
667/// were not (the upload starts within the VRAM budget, so a later call
668/// finds them warm) or the tensor cannot go to the GPU at all. Keeps the
669/// probe from billing a full cold dispatch+readback to a sample it will
670/// discard anyway. The verdict needs only a couple of warm tensors, so
671/// probe-driven uploads are capped — the losing-GPU machine should not
672/// pay for uploading the whole layer stack it will never use; if the GPU
673/// wins, the rest uploads lazily on demand, in the same first-touch order.
674#[allow(unused_variables)]
675pub fn q8_resident_or_upload(model: &Arc<CmfModel>, idx: usize) -> bool {
676    static PROBE_UPLOADS: AtomicU32 = AtomicU32::new(0);
677    let may_upload = PROBE_UPLOADS.load(Ordering::Relaxed) < 4;
678    let resident = match backend() {
679        #[cfg(target_os = "macos")]
680        Backend::Metal => crate::gpu_metal::q8_resident_or_upload(model, idx, may_upload),
681        #[cfg(feature = "gpu")]
682        Backend::Wgpu => crate::gpu_wgpu::q8_resident_or_upload(model, idx, may_upload),
683        Backend::None => false,
684    };
685    if !resident && may_upload {
686        PROBE_UPLOADS.fetch_add(1, Ordering::Relaxed);
687    }
688    resident
689}
690
691/// Test hook: reset all probes to the undecided state.
692#[cfg(test)]
693pub(crate) fn probe_reset() {
694    for p in &PROBES {
695        p.state.store(0, Ordering::Relaxed);
696        p.flip.store(0, Ordering::Relaxed);
697        p.gpu_ns.store(0, Ordering::Relaxed);
698        p.gpu_n.store(0, Ordering::Relaxed);
699        p.cpu_ns.store(0, Ordering::Relaxed);
700        p.cpu_n.store(0, Ordering::Relaxed);
701    }
702}
703
704#[cfg(test)]
705mod probe_tests {
706    use super::*;
707    use std::time::Duration;
708
709    // One test fn: PROBES is process-global and probe_reset touches all
710    // classes — parallel test threads would race.
711    #[test]
712    fn probe_alternates_discards_cold_and_decides() {
713        probe_reset();
714        // Probing: arms alternate.
715        assert!(matches!(probe_arm(OpClass::Ffn), ProbeArm::Gpu));
716        assert!(matches!(probe_arm(OpClass::Ffn), ProbeArm::CpuTimed));
717
718        // A cold GPU sample (upload noted) must be discarded: feed a
719        // catastrophic cold sample, then clean fast-GPU samples — GPU
720        // wins only if the cold one did not count.
721        probe_note_cold();
722        probe_record(OpClass::Ffn, true, Duration::from_secs(1000));
723        for _ in 0..PROBE_SAMPLES {
724            probe_record(OpClass::Ffn, true, Duration::from_millis(1));
725            probe_record(OpClass::Ffn, false, Duration::from_millis(4));
726        }
727        assert!(matches!(probe_arm(OpClass::Ffn), ProbeArm::Gpu));
728
729        // The reverse: a class where the CPU arm is faster decides CPU.
730        for _ in 0..PROBE_SAMPLES {
731            probe_record(OpClass::Matmat, true, Duration::from_millis(4));
732            probe_record(OpClass::Matmat, false, Duration::from_millis(1));
733        }
734        assert!(matches!(probe_arm(OpClass::Matmat), ProbeArm::Cpu));
735
736        // cpu_scope: gates off inside, restored after.
737        cpu_scope(|| CPU_ONLY.with(|c| assert!(c.get())));
738        CPU_ONLY.with(|c| assert!(!c.get()));
739        cpu_scope(|| {
740            cpu_scope(|| CPU_ONLY.with(|c| assert!(c.get())));
741            CPU_ONLY.with(|c| assert!(c.get()));
742        });
743        let _ = std::panic::catch_unwind(|| cpu_scope(|| panic!("scope test")));
744        CPU_ONLY.with(|c| assert!(!c.get()));
745        probe_reset();
746    }
747
748    #[test]
749    fn a_remembered_verdict_is_adopted_and_a_stranger_is_not() {
750        // Probing is not free: on a Snapdragon 778G the deciding classes
751        // cost minutes of wall clock before the first token, every
752        // process, and reached the same verdict every time. The cache
753        // exists so that price is paid once.
754        //
755        // The key is built from THIS process's device, never a name this
756        // test sets: `probe_set_device` is first-writer-wins and on a Mac
757        // the Metal backend may already have named the silicon before the
758        // tests run — which is exactly how this test failed on CI while
759        // passing locally. GemmNt on purpose: the arbitration test never
760        // touches it, and both run in one process.
761        let mine = probe_cache_key_named("gemm-nt");
762        let state = || {
763            PROBES[OpClass::GemmNt as usize]
764                .state
765                .load(Ordering::Relaxed)
766        };
767
768        // Another device's verdict is not mine, whatever it claims.
769        probe_cache_adopt("SomeOtherGPU/Vulkan\tgemm-nt\tgpu\n");
770        assert_eq!(state(), 0);
771        // Neither is one from another build of this engine.
772        let older = mine.replacen(env!("CARGO_PKG_VERSION"), "0.0.0-old", 1);
773        assert_ne!(older, mine);
774        probe_cache_adopt(&format!("{older}\tgpu\n"));
775        assert_eq!(state(), 0);
776        // Mine is.
777        probe_cache_adopt(&format!("{mine}\tcpu\n"));
778        assert_eq!(state(), 2);
779
780        PROBES[OpClass::GemmNt as usize]
781            .state
782            .store(0, Ordering::Relaxed);
783    }
784}
785
786/// Default row threshold: the GPU takes only larger matrices (lm_head
787/// class). Below it, the dispatch/readback cost does not pay off on unified memory.
788pub const GPU_MIN_ROWS: usize = 65_536;
789
790/// Effective threshold: `CMF_GPU_MIN_ROWS` overrides. Defaults differ
791/// by device class: on a DISCRETE card VRAM bandwidth pays off even for
792/// FFN/QKV-class matrices (4096), on unified memory only lm_head-class
793/// is worth the dispatch/readback (65536). Field case behind this: a
794/// 35B model on an RTX 4090 saw ~0 offload because every layer matrix
795/// sat below the old universal 65536.
796pub fn min_rows() -> usize {
797    if let Some(v) = std::env::var("CMF_GPU_MIN_ROWS")
798        .ok()
799        .and_then(|v| v.parse().ok())
800    {
801        return v;
802    }
803    if discrete() { 4096 } else { GPU_MIN_ROWS }
804}
805
806/// Is the active backend a discrete card (PCIe VRAM)?
807pub fn discrete() -> bool {
808    match backend() {
809        #[cfg(feature = "gpu")]
810        Backend::Wgpu => crate::gpu_wgpu::is_discrete(),
811        #[cfg(target_os = "macos")]
812        Backend::Metal => false, // UMA by the init() guard
813        Backend::None => false,
814    }
815}
816
817/// A single MoE-FFN job (an expert with its own weight), executed in one
818/// submission: (rows, cols, idx, row_scale) for gate/up/down + prescaled
819/// inputs + the down θ-field + the blending weight.
820pub struct MoeJob<'a> {
821    pub gate: (usize, usize, usize, &'a [f32]),
822    pub up: (usize, usize, usize, &'a [f32]),
823    pub down: (usize, usize, usize, &'a [f32]),
824    pub xs_gate: Vec<f32>,
825    pub xs_up: Vec<f32>,
826    pub down_col: &'a [f32],
827    pub w: f32,
828    /// q1 trio: scales live inside the 6-byte tiles (row_scale slices
829    /// empty, xs raw f32). Backends without a q1 kernel refuse the job.
830    pub q1: bool,
831    /// q4_tiled trio: scales inside the 18-byte tiles (row_scale
832    /// slices empty, xs raw f32) — the MoE-hybrid coder class.
833    pub q4t: bool,
834    /// q4tp trio: same raw-xs contract, 16-byte nibble stride and the scale
835    /// on a per-row ladder. Without this the experts of a q4tp MoE model fall
836    /// to the CPU while every other dtype rides the device.
837    pub q4tp: bool,
838    /// Mixed 2-bit profile: gate/up are q2tp (8-byte chunks, zero rung),
839    /// down stays q4tp. Set together with `q4tp`; a backend without the
840    /// 2-bit kernel must refuse the whole job.
841    pub gu_q2: bool,
842    /// The reference's `swiglu_limit`; 0 disables the clamp. A backend that
843    /// cannot apply it must REFUSE the job rather than drop it silently —
844    /// the difference only shows on saturating activations, which is the
845    /// hardest kind of divergence to notice.
846    pub swiglu_limit: f32,
847}
848
849/// A single independent batch matvec (GDN projections of one input).
850pub struct BatchJob<'a> {
851    pub idx: usize,
852    pub rows: usize,
853    pub cols: usize,
854    pub row_scale: &'a [f32],
855    pub xs: Vec<f32>,
856    /// Weight layout. Was a bare `q1: bool`, which could only ever spell two
857    /// of the four and silently sent everything else back to the CPU — the
858    /// GDN projections of a q4t/q4tp model never reached the device at all.
859    pub layout: BatchLayout,
860}
861
862/// Which kernel a batched matvec needs. q8 carries row scales in a side
863/// buffer; the rest embed them in the payload and differ in stride.
864#[derive(Clone, Copy, PartialEq, Eq, Debug)]
865pub enum BatchLayout {
866    Q8,
867    Q1,
868    Q4t,
869    Q4tp,
870}
871
872#[derive(Clone, Copy, PartialEq, Eq)]
873enum Backend {
874    None,
875    #[cfg(target_os = "macos")]
876    Metal,
877    #[cfg(feature = "gpu")]
878    Wgpu,
879}
880
881fn backend() -> Backend {
882    #[cfg(feature = "gpu")]
883    if crate::gpu_wgpu::selected() {
884        return if crate::gpu_wgpu::enabled() {
885            Backend::Wgpu
886        } else {
887            Backend::None
888        };
889    }
890    #[cfg(target_os = "macos")]
891    if crate::gpu_metal::enabled() {
892        return Backend::Metal;
893    }
894    Backend::None
895}
896
897/// GPU enabled and initialized on the selected backend?
898/// Whether THIS build can bring a GPU up on THIS device: a compiled-in
899/// backend plus a live adapter. The mobile FFI exposes it so an app can
900/// tell "GPU off" from "GPU impossible" (a CPU-only .so ships no
901/// backend at all). Cached after the first call.
902pub fn backend_available() -> bool {
903    #[cfg(target_os = "macos")]
904    {
905        // The Metal path is always compiled on macOS.
906        true
907    }
908    #[cfg(all(feature = "gpu", not(target_os = "macos")))]
909    {
910        static AVAIL: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
911        *AVAIL.get_or_init(crate::gpu_wgpu::adapter_probe)
912    }
913    #[cfg(all(not(feature = "gpu"), not(target_os = "macos")))]
914    {
915        false
916    }
917}
918
919/// A process-wide, phase-scoped GPU gate. `cpu_scope` is thread-local and
920/// the pool's workers do not inherit it, so a caller that wants a whole
921/// *phase* off the device — a prompt encoder whose weights live in a part of
922/// the file the hot loop never touches, on a machine that cannot keep both
923/// wired — has to say so globally.
924static GPU_PAUSED: AtomicBool = AtomicBool::new(false);
925
926/// Park the device for every thread until the returned guard drops.
927pub fn pause_gpu() -> GpuPause {
928    GPU_PAUSED.store(true, Ordering::Relaxed);
929    GpuPause(())
930}
931
932pub struct GpuPause(());
933
934impl Drop for GpuPause {
935    fn drop(&mut self) {
936        GPU_PAUSED.store(false, Ordering::Relaxed);
937    }
938}
939
940pub fn enabled() -> bool {
941    !GPU_PAUSED.load(Ordering::Relaxed) && backend() != Backend::None
942}
943
944/// Default-on condition for the wgpu whole-token graph: the wgpu
945/// backend on a DISCRETE adapter. NOT plain `enabled()` (macOS/Metal
946/// must not pay a per-token layer scan for a graph its backend
947/// refuses), and NOT integrated adapters: the graph's ~300 barriered
948/// dispatches per token are cheap on desktop immediate-mode GPUs but
949/// tiled mobile GPUs (Adreno/Mali) drain the pipeline at every barrier
950/// — field report: 0.2 tok/s on-graph vs 15 tok/s on the CPU. On
951/// integrated adapters the per-op probe path arbitrates each op class
952/// against the CPU instead; CMF_GPU_WGPU_GRAPH=1 still forces the
953/// graph anywhere.
954/// Is the wgpu backend active at all (any adapter)? Eligibility gate
955/// for the whole-token graph — whether it actually RUNS is decided by
956/// `wgpu_graph_default` (trusted on discrete) or the generation race.
957pub fn wgpu_active() -> bool {
958    #[cfg(feature = "gpu")]
959    {
960        matches!(backend(), Backend::Wgpu)
961    }
962    #[cfg(not(feature = "gpu"))]
963    {
964        false
965    }
966}
967
968/// Which GPU this thread's engine calls address. Multi-card hosts hold
969/// one wgpu context PER card (weights, KV mirrors and scratch live
970/// inside a context, so per-device contexts give per-device caches for
971/// free); this thread-local says which one is current. Default: the
972/// process pin (CMF_GPU_ADAPTER) or 0 — so single-card runs behave
973/// exactly as they always have.
974pub fn default_device() -> usize {
975    static D: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
976    *D.get_or_init(|| {
977        std::env::var("CMF_GPU_ADAPTER")
978            .ok()
979            .and_then(|v| v.trim().parse::<usize>().ok())
980            .unwrap_or(0)
981    })
982}
983
984thread_local! {
985    static CUR_DEV: std::cell::Cell<Option<usize>> = const { std::cell::Cell::new(None) };
986}
987
988/// The device this thread is pinned to.
989pub fn current_device() -> usize {
990    CUR_DEV.with(|c| c.get()).unwrap_or_else(default_device)
991}
992
993/// Pin this thread to a device. Server slots call it once per request;
994/// the worker pool propagates it into its threads, so a dispatch begun
995/// on card 1 does not finish on card 0.
996pub fn set_current_device(i: usize) {
997    CUR_DEV.with(|c| c.set(Some(i)));
998}
999
1000/// Run `f` with this thread pinned to `dev`, restoring the previous pin.
1001pub fn with_device<R>(dev: usize, f: impl FnOnce() -> R) -> R {
1002    let prev = CUR_DEV.with(|c| c.replace(Some(dev)));
1003    let r = f();
1004    CUR_DEV.with(|c| c.set(prev));
1005    r
1006}
1007
1008/// How many GPUs this process can address (wgpu adapter count; 1 on
1009/// Metal, 0 without a backend).
1010pub fn device_count() -> usize {
1011    #[cfg(all(feature = "gpu", not(target_os = "macos")))]
1012    {
1013        return crate::gpu_wgpu::adapter_count();
1014    }
1015    #[cfg(not(all(feature = "gpu", not(target_os = "macos"))))]
1016    {
1017        usize::from(backend_available())
1018    }
1019}
1020
1021/// Weight budget of the current GPU in bytes; 0 when there is none and
1022/// u64::MAX on unified memory (where the OS pages shared RAM and the
1023/// question "does the model fit the card" has no separate answer).
1024pub fn vram_budget() -> u64 {
1025    #[cfg(all(feature = "gpu", not(target_os = "macos")))]
1026    {
1027        return crate::gpu_wgpu::device_vram_budget();
1028    }
1029    #[cfg(not(all(feature = "gpu", not(target_os = "macos"))))]
1030    {
1031        if backend_available() { u64::MAX } else { 0 }
1032    }
1033}
1034
1035/// Device weight bytes uploaded so far (wgpu; 0 on other backends).
1036/// Steady-state windows must show a ZERO delta — growth mid-benchmark
1037/// means eviction/re-upload and disqualifies the number.
1038pub fn upload_bytes() -> u64 {
1039    #[cfg(feature = "gpu")]
1040    {
1041        return crate::gpu_wgpu::UPLOAD_BYTES.load(std::sync::atomic::Ordering::Relaxed);
1042    }
1043    #[cfg(not(feature = "gpu"))]
1044    0
1045}
1046
1047/// Which half of the run is asking.
1048///
1049/// The phase exists because the graph is plausibly two decisions, not
1050/// one — but on the hardware measured so far it is only ever a decode
1051/// decision. On an Adreno 642L with bonsai-1.7b, from identical clean
1052/// starts and two repeats each: decode 11.6 tok/s without it and 0.72
1053/// with, while prefill is 4.2 either way. A first reading of 3.4 -> 18.0
1054/// for prefill did not survive a controlled re-run — it was a dirty
1055/// probe cache between configurations, not the graph, and the prefill
1056/// route through the graph is GDN-only in the first place, which this
1057/// dense model never takes.
1058#[derive(Clone, Copy, PartialEq, Eq, Debug)]
1059pub enum GraphPhase {
1060    Prefill,
1061    Decode,
1062}
1063
1064/// The one place that decides whether the whole-token graph runs.
1065///
1066/// `CMF_GPU_WGPU_GRAPH`: `0` off everywhere, `prefill` only for the
1067/// prompt, anything else on everywhere. Unset: desktop-class GPUs take
1068/// it for both phases; phone-class UMA takes it for PREFILL only, which
1069/// is the measurement above rather than a guess — the per-op path keeps
1070/// decode, where it is seventeen times better.
1071pub fn wgpu_graph_on(phase: GraphPhase) -> bool {
1072    match std::env::var("CMF_GPU_WGPU_GRAPH").ok().as_deref() {
1073        Some("0") => false,
1074        Some("prefill") => phase == GraphPhase::Prefill,
1075        Some(_) => true,
1076        None => {
1077            if wgpu_graph_default() {
1078                return true;
1079            }
1080            // Integrated/mobile keeps the per-op path for BOTH phases —
1081            // unchanged, because the measurement that would have bought
1082            // prefill a graph did not reproduce. `=prefill` is there for
1083            // the device where it does; the default does not guess.
1084            let _ = phase;
1085            false
1086        }
1087    }
1088}
1089
1090pub fn wgpu_graph_default() -> bool {
1091    #[cfg(feature = "gpu")]
1092    {
1093        // Discrete cards always; Apple-silicon UMA on macOS too — desktop
1094        // -class GPUs where the graph measured ~2x the CPU on the Qwen3.6
1095        // family (M4: 13.3 tok/s against 7.3). Phone-class UMA (Android/
1096        // iOS builds) keeps the per-op probe path: tiled mobile GPUs have
1097        // turned the ~300-dispatch graph into seconds per token.
1098        matches!(backend(), Backend::Wgpu)
1099            && (crate::gpu_wgpu::discrete_active()
1100                || (cfg!(target_os = "macos") && crate::gpu_wgpu::adapter_up()))
1101    }
1102    #[cfg(not(feature = "gpu"))]
1103    {
1104        false
1105    }
1106}
1107
1108/// q8_row/q8_2f matvec, rows [row0, row0+rows). `xs` — prescaled by the θ-field.
1109#[allow(clippy::too_many_arguments, unused_variables)]
1110pub fn q8_matvec_range(
1111    model: &Arc<CmfModel>,
1112    idx: usize,
1113    row0: usize,
1114    row_scale: &[f32],
1115    xs: &[f32],
1116    rows: usize,
1117    cols: usize,
1118    out: &mut [f32],
1119) -> bool {
1120    match backend() {
1121        #[cfg(target_os = "macos")]
1122        Backend::Metal => {
1123            crate::gpu_metal::q8_matvec_range(model, idx, row0, row_scale, xs, rows, cols, out)
1124        }
1125        #[cfg(feature = "gpu")]
1126        Backend::Wgpu => {
1127            crate::gpu_wgpu::q8_matvec_range(model, idx, row0, row_scale, xs, rows, cols, out)
1128        }
1129        Backend::None => false,
1130    }
1131}
1132
1133/// GEMM of a prefill batch: `pre` — prescaled inputs row-major [b, cols],
1134/// out — row-major [b, rows].
1135#[allow(clippy::too_many_arguments, unused_variables)]
1136/// The two-field int8 GEMM with the column field left for the device.
1137/// wgpu only — Metal's int8 kernel takes a pre-scaled activation, so the
1138/// caller keeps that path when this returns `false`.
1139#[allow(clippy::too_many_arguments)]
1140pub fn q8_matmat_2f(
1141    model: &Arc<CmfModel>,
1142    idx: usize,
1143    row_scale: &[f32],
1144    col_field: &[f32],
1145    xs: &[f32],
1146    b: usize,
1147    rows: usize,
1148    cols: usize,
1149    out: &mut [f32],
1150) -> bool {
1151    #[allow(unreachable_patterns)]
1152    match backend() {
1153        #[cfg(feature = "gpu")]
1154        Backend::Wgpu => {
1155            crate::gpu_wgpu::q8_matmat_2f(model, idx, row_scale, col_field, xs, b, rows, cols, out)
1156        }
1157        _ => false,
1158    }
1159}
1160
1161pub fn q8_matmat(
1162    model: &Arc<CmfModel>,
1163    idx: usize,
1164    row_scale: &[f32],
1165    pre: &[f32],
1166    b: usize,
1167    rows: usize,
1168    cols: usize,
1169    out: &mut [f32],
1170) -> bool {
1171    match backend() {
1172        #[cfg(target_os = "macos")]
1173        Backend::Metal => {
1174            crate::gpu_metal::q8_matmat(model, idx, row_scale, pre, b, rows, cols, out)
1175        }
1176        #[cfg(feature = "gpu")]
1177        Backend::Wgpu => crate::gpu_wgpu::q8_matmat(model, idx, row_scale, pre, b, rows, cols, out),
1178        Backend::None => false,
1179    }
1180}
1181
1182/// q1 matvec: raw f32 activations, tile-embedded scales. Metal only
1183/// for now (wgpu q1 WGSL is queued); false = CPU fallback.
1184#[allow(unused_variables)]
1185pub fn q1_matvec(
1186    model: &Arc<CmfModel>,
1187    idx: usize,
1188    xs: &[f32],
1189    rows: usize,
1190    cols: usize,
1191    out: &mut [f32],
1192) -> bool {
1193    match backend() {
1194        #[cfg(target_os = "macos")]
1195        Backend::Metal => crate::gpu_metal::q1_matvec(model, idx, xs, rows, cols, out),
1196        #[cfg(feature = "gpu")]
1197        Backend::Wgpu => crate::gpu_wgpu::q1_matvec(model, idx, xs, rows, cols, out),
1198        Backend::None => false,
1199    }
1200}
1201
1202/// Whole attention sub-block on the wgpu token graph (drop-in for
1203/// `qwen_attention`): normed hidden in, O-projection out, resident device
1204/// K/V mirror. false = refusal / not the wgpu backend → CPU path.
1205#[allow(clippy::too_many_arguments)]
1206pub fn attn_dropin(
1207    model: &Arc<CmfModel>,
1208    kv_id: u64,
1209    layer: usize,
1210    normed: &[f32],
1211    wq_idx: usize,
1212    wk_idx: usize,
1213    wv_idx: usize,
1214    wo_idx: usize,
1215    q_norm: Option<&[f32]>,
1216    k_norm: Option<&[f32]>,
1217    invf: &[f32],
1218    nh: usize,
1219    nkv: usize,
1220    hd: usize,
1221    rd: usize,
1222    hidden: usize,
1223    pos: usize,
1224    cap: usize,
1225    gemma: bool,
1226    eps: f32,
1227    cpu_k: &[Vec<f32>],
1228    cpu_v: &[Vec<f32>],
1229    out: &mut [f32],
1230) -> bool {
1231    match backend() {
1232        #[cfg(feature = "gpu")]
1233        Backend::Wgpu => crate::gpu_wgpu::attn_dropin_gpu(
1234            model, kv_id, layer, normed, wq_idx, wk_idx, wv_idx, wo_idx, q_norm, k_norm, invf, nh,
1235            nkv, hd, rd, hidden, pos, cap, gemma, eps, cpu_k, cpu_v, out,
1236        ),
1237        #[allow(unused_variables)]
1238        _ => false,
1239    }
1240}
1241
1242/// One weight in the whole-token graph: tensor idx + a codec tag (0=q8_row,
1243/// 1=q1, 2=q4_tiled, 3=q1t, 4=f32) + per-row scales (q8_row only) + the raw f32
1244/// data (kind 4 only — small unquantized projections like GDN in_proj_a/b).
1245pub struct GraphW<'a> {
1246    pub idx: usize,
1247    pub kind: u8,
1248    pub row_scale: &'a [f32],
1249    pub data: &'a [f32],
1250}
1251
1252/// A layer's token-mixing op: standard attention or a GDN (linear-attention)
1253/// block. The surrounding norms + SwiGLU FFN are common to both.
1254pub enum GraphAttn<'a> {
1255    Full {
1256        wq: GraphW<'a>,
1257        wk: GraphW<'a>,
1258        wv: GraphW<'a>,
1259        wo: GraphW<'a>,
1260        q_norm: Option<&'a [f32]>,
1261        k_norm: Option<&'a [f32]>,
1262        /// (bq, bk, bv) attention biases (Qwen2). None ⇒ no bias.
1263        bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
1264        /// Qwen3.5 gated attention: wq emits 2·nh·hd (q||gate per head), the
1265        /// attention output is scaled by sigmoid(gate) before the O projection.
1266        output_gate: bool,
1267        cpu_k: &'a [Vec<f32>],
1268        cpu_v: &'a [Vec<f32>],
1269    },
1270    Gdn {
1271        qkv: GraphW<'a>,
1272        z: GraphW<'a>,
1273        a: GraphW<'a>,
1274        b: GraphW<'a>,
1275        out: GraphW<'a>,
1276        conv1d: &'a [f32],
1277        a_log: &'a [f32],
1278        dt_bias: &'a [f32],
1279        norm: &'a [f32],
1280        nv: usize,
1281        nk: usize,
1282        dk: usize,
1283        dv: usize,
1284        kk: usize,
1285        /// CPU recurrent state `[ring (kk-1)·cdim | S nv·dk·dv]` — seeds the
1286        /// device mirror when prefill ran on the host (o1 collection, CPU
1287        /// fallback): a zero-initialized device state at decode is exactly
1288        /// the "coherent but contextless" garble.
1289        cpu_state: &'a [f32],
1290    },
1291    /// LFM2 gated short convolution: a fused (B, C, x) projection, a
1292    /// depthwise causal conv over a (kernel−1)-deep per-channel ring,
1293    /// C-gating, and an output projection. This mixer is what most of an
1294    /// LFM2 stack is (22 of the 2.6B's 30 layers), and before it had a
1295    /// graph arm the whole model fell to the per-op path — ~100 submits
1296    /// a token, 22 tok/s on an A100 for a 1.4 GB file.
1297    ShortConv {
1298        /// [3·hidden, hidden] fused input projection.
1299        inp: GraphW<'a>,
1300        /// [hidden, hidden] output projection.
1301        out: GraphW<'a>,
1302        /// [hidden · kernel] depthwise taps, `[channel][tap]`, tap
1303        /// kernel−1 multiplying the current position.
1304        taps: &'a [f32],
1305        kernel: usize,
1306        /// CPU conv ring `[channel][kernel−1]`, slot 0 newest — seeds
1307        /// the device mirror when prefill ran on the host, which for
1308        /// this mixer is always (the batch graph declines it).
1309        cpu_state: &'a [f32],
1310    },
1311}
1312
1313/// Per-layer weights for the whole-token wgpu graph.
1314pub struct GraphLayer<'a> {
1315    pub input_norm: &'a [f32],
1316    pub attn: GraphAttn<'a>,
1317    pub post_norm: &'a [f32],
1318    pub ffn: GraphFfn<'a>,
1319}
1320
1321/// The FFN of one graph layer: a dense SwiGLU trio, or a routed MoE —
1322/// router + top-k selection + all selected experts run ON DEVICE (the
1323/// routing decision depends on the resident hidden state, so a CPU
1324/// round-trip per layer would forfeit the one-submit design).
1325pub enum GraphFfn<'a> {
1326    Dense {
1327        gate: GraphW<'a>,
1328        up: GraphW<'a>,
1329        down: GraphW<'a>,
1330    },
1331    Moe {
1332        /// Router logits weight (f32, kind 4) `[n_exp, hidden]`.
1333        router: GraphW<'a>,
1334        /// Shared-expert sigmoid gate (f32) `[1, hidden]`.
1335        shared_gate: GraphW<'a>,
1336        /// Per-expert q4_tiled directory indices `(gate, up, down)`;
1337        /// the SHARED expert rides as the LAST entry — the select
1338        /// kernel pins it with the sigmoid weight.
1339        experts: Vec<(usize, usize, usize)>,
1340        /// Routed experts (shared excluded).
1341        n_exp: usize,
1342        top_k: usize,
1343        inter: usize,
1344        norm_topk: bool,
1345        /// Expert weight layout, uniform across the layer: `false` =
1346        /// q4_tiled (18 B tiles, inline f16 scale), `true` = q4tp
1347        /// (16 B nibbles + a per-row ladder plane). The two differ only
1348        /// in where the scale comes from, so they share every kernel
1349        /// but the weight-staging block.
1350        q4tp: bool,
1351        /// `true` = the gate/up experts are `q2tp` (2-bit plane) while
1352        /// `down` stays q4tp — the mixed profile a 2-bit-class checkpoint
1353        /// converts into. Only meaningful with `q4tp: true`.
1354        gu_q2: bool,
1355        /// LFM2-MoE / DeepSeek-V3 `noaux_tc` routing: per-expert sigmoid
1356        /// scores instead of a softmax, and `norm_topk` renormalises with
1357        /// the 1e-6 floor. The softmax arm is bit-identical to before.
1358        sigmoid: bool,
1359        /// Per-expert SELECTION bias: added to the score for the top-k
1360        /// choice only — the mixing weights stay unbiased (noaux_tc).
1361        bias: Option<&'a [f32]>,
1362        /// Whether a shared expert rides as the last `experts` entry.
1363        /// LFM2-MoE has none; the select kernel then leaves slot `top_k`
1364        /// unwritten and the expert loop runs `top_k` slots, not +1.
1365        has_shared: bool,
1366    },
1367}
1368
1369/// Whole-token decode graph on wgpu: the entire layer stack in ONE submit,
1370/// hidden resident, one readback. Updates `h` in place. false = refusal.
1371/// `loop_norm_at`: virtual layer indices after which `final_norm` is applied
1372/// (Looped Transformer mid-stack norm). Empty for standard models.
1373#[allow(clippy::too_many_arguments)]
1374pub fn forward_token_graph(
1375    model: &Arc<CmfModel>,
1376    kv_id: u64,
1377    layers: &[GraphLayer],
1378    // Per-layer sealed o1 (Nystrom) state; Some = replace this layer's
1379    // exact attention with the O(1) kernels. wgpu only.
1380    o1: &[Option<Vec<crate::nystrom::O1DeviceView<'_>>>],
1381    o1_epoch: u64,
1382    invf: &[f32],
1383    h: &mut [f32],
1384    nh: usize,
1385    nkv: usize,
1386    hd: usize,
1387    rd: usize,
1388    hidden: usize,
1389    inter: usize,
1390    position: usize,
1391    cap: usize,
1392    gemma: bool,
1393    eps: f32,
1394    lm_head: Option<(&GraphW, usize)>,
1395    final_norm: &[f32],
1396    logits: &mut Vec<f32>,
1397    loop_norm_at: &[usize],
1398    steps: usize,
1399    embed: Option<(&GraphW, usize, f32)>,
1400    ids_out: Option<&mut Vec<u32>>,
1401    // How many leading layers the graph ran (see the wgpu twin) — smaller
1402    // than layers.len() when the expert budget ended the device prefix.
1403    layers_run: Option<&mut usize>,
1404    // Absolute index of layers[0] in the model — the KV/state mirrors key
1405    // on it, so a layer SPAN (network split segment) shares mirrors with
1406    // a full-stack run instead of colliding at slot 0.
1407    layer_base: usize,
1408    // Read the final hidden back alongside the fused head's logits.
1409    hidden_too: bool,
1410) -> bool {
1411    match backend() {
1412        #[cfg(feature = "gpu")]
1413        Backend::Wgpu => crate::gpu_wgpu::forward_token_graph(
1414            model,
1415            kv_id,
1416            layers,
1417            o1,
1418            o1_epoch,
1419            invf,
1420            h,
1421            nh,
1422            nkv,
1423            hd,
1424            rd,
1425            hidden,
1426            inter,
1427            position,
1428            cap,
1429            gemma,
1430            eps,
1431            lm_head,
1432            final_norm,
1433            logits,
1434            loop_norm_at,
1435            steps,
1436            embed,
1437            ids_out,
1438            layers_run,
1439            layer_base,
1440            hidden_too,
1441        ),
1442        #[allow(unused_variables)]
1443        _ => {
1444            let _ = (
1445                lm_head,
1446                final_norm,
1447                logits,
1448                loop_norm_at,
1449                layers_run,
1450                layer_base,
1451                hidden_too,
1452            );
1453            false
1454        }
1455    }
1456}
1457
1458/// Speculative-verify tail for the batched graph: fold final-norm + lm_head
1459/// over every batch position and read all k logit rows back; the batch also
1460/// snapshots the GDN state per position for `gdn_spec_restore`.
1461pub struct SpecTail<'a> {
1462    pub lm: GraphW<'a>,
1463    pub lm_rows: usize,
1464    pub final_norm: &'a [f32],
1465    pub logits_out: &'a mut Vec<f32>,
1466}
1467
1468/// Batched prefill: k contiguous positions through the whole graph in one submit
1469/// (projections/FFN as GEMMs, attention/GDN looped over scratch). `h` is
1470/// [k·hidden] in/out; `positions` len k. wgpu only.
1471#[allow(clippy::too_many_arguments)]
1472pub fn forward_batch_graph(
1473    model: &Arc<CmfModel>,
1474    kv_id: u64,
1475    layers: &[GraphLayer],
1476    invf: &[f32],
1477    h: &mut [f32],
1478    nh: usize,
1479    nkv: usize,
1480    hd: usize,
1481    rd: usize,
1482    hidden: usize,
1483    inter: usize,
1484    positions: &[usize],
1485    cap: usize,
1486    gemma: bool,
1487    eps: f32,
1488    k: usize,
1489    spec: Option<SpecTail<'_>>,
1490) -> bool {
1491    match backend() {
1492        #[cfg(feature = "gpu")]
1493        Backend::Wgpu => crate::gpu_wgpu::forward_batch_graph(
1494            model, kv_id, layers, invf, h, nh, nkv, hd, rd, hidden, inter, positions, cap, gemma,
1495            eps, k, spec,
1496        ),
1497        #[allow(unreachable_patterns)]
1498        _ => {
1499            let _ = spec;
1500            false
1501        }
1502    }
1503}
1504
1505/// After a partial speculative acceptance: restore every GDN layer's device
1506/// state to the snapshot after batch position `slot`. wgpu only.
1507pub fn gdn_spec_restore(kv_id: u64, slot: usize) -> bool {
1508    #[cfg(feature = "gpu")]
1509    if backend() == Backend::Wgpu {
1510        return crate::gpu_wgpu::gdn_spec_restore(kv_id, slot);
1511    }
1512    #[allow(unreachable_code)]
1513    {
1514        let _ = (kv_id, slot);
1515        false
1516    }
1517}
1518
1519/// Drop the wgpu token graph's device K/V mirror for a pipeline.
1520pub fn graph_kv_reset(_kv_id: u64) {
1521    #[cfg(feature = "gpu")]
1522    if backend() == Backend::Wgpu {
1523        crate::gpu_wgpu::kv_mirror_reset(_kv_id);
1524    }
1525}
1526
1527/// Ternary (q1t) BASE matvec on the GPU — fills `out` with the base dot; the
1528/// caller adds the sparse overlay on the CPU. Metal only for now (wgpu q1t not
1529/// yet written → CPU fallback).
1530pub fn q1t_matvec(
1531    model: &Arc<CmfModel>,
1532    idx: usize,
1533    xs: &[f32],
1534    rows: usize,
1535    cols: usize,
1536    out: &mut [f32],
1537) -> bool {
1538    match backend() {
1539        #[cfg(target_os = "macos")]
1540        Backend::Metal => {
1541            if metal_q1t_enabled() {
1542                crate::gpu_metal::q1t_matvec(model, idx, xs, rows, cols, out)
1543            } else {
1544                false
1545            }
1546        }
1547        #[cfg(feature = "gpu")]
1548        Backend::Wgpu => crate::gpu_wgpu::q1t_matvec(model, idx, xs, rows, cols, out),
1549        Backend::None => false,
1550    }
1551}
1552
1553/// q4_block matvec on the GPU — wgpu only (Metal drives q4_block through the
1554/// whole-token graph, not a standalone matvec).
1555#[allow(unused_variables)]
1556pub fn q4b_matvec(
1557    model: &Arc<CmfModel>,
1558    idx: usize,
1559    xs: &[f32],
1560    rows: usize,
1561    cols: usize,
1562    out: &mut [f32],
1563) -> bool {
1564    match backend() {
1565        #[cfg(target_os = "macos")]
1566        Backend::Metal => false,
1567        #[cfg(feature = "gpu")]
1568        Backend::Wgpu => crate::gpu_wgpu::q4b_matvec(model, idx, xs, rows, cols, out),
1569        Backend::None => false,
1570    }
1571}
1572
1573/// q1t batched GEMM (prefill) — base + overlay on-device (Metal simdgroup or
1574/// wgpu register-blocked).
1575pub fn q1t_matmat(
1576    model: &Arc<CmfModel>,
1577    idx: usize,
1578    xs: &[f32],
1579    b: usize,
1580    rows: usize,
1581    cols: usize,
1582    out: &mut [f32],
1583) -> bool {
1584    match backend() {
1585        #[cfg(target_os = "macos")]
1586        // Batched prefill and single-token decode are both enabled. On the
1587        // real 14.8B Q1T model prefill PPL was within 0.3% of CPU (7.942 vs
1588        // 7.966), and the alignment-safe decode kernel reached 3.52e-6 max_rel.
1589        Backend::Metal => crate::gpu_metal::q1t_matmat(model, idx, xs, b, rows, cols, out),
1590        #[cfg(feature = "gpu")]
1591        Backend::Wgpu => crate::gpu_wgpu::q1t_matmat(model, idx, xs, b, rows, cols, out),
1592        Backend::None => false,
1593    }
1594}
1595
1596/// Native Metal Q1T switch. Enabled by default after the byte-packed Q1T
1597/// fields were changed to alignment-safe loads; keep an explicit emergency
1598/// fallback for device/driver diagnostics.
1599#[cfg(target_os = "macos")]
1600pub(crate) fn metal_q1t_enabled() -> bool {
1601    std::env::var("CMF_METAL_Q1T")
1602        .map(|v| v != "0" && !v.eq_ignore_ascii_case("off"))
1603        .unwrap_or(true)
1604}
1605
1606/// Batched q1 GEMM (prefill). wgpu only — Metal has its own block path.
1607pub fn q1_matmat(
1608    model: &Arc<CmfModel>,
1609    idx: usize,
1610    xs: &[f32],
1611    b: usize,
1612    rows: usize,
1613    cols: usize,
1614    out: &mut [f32],
1615) -> bool {
1616    match backend() {
1617        #[cfg(feature = "gpu")]
1618        Backend::Wgpu => crate::gpu_wgpu::q1_matmat(model, idx, xs, b, rows, cols, out),
1619        #[allow(unused_variables)]
1620        _ => false,
1621    }
1622}
1623
1624/// Contention kill for the wide imagegen GEMM/FFN paths: one grossly
1625/// slow op under a work-proportional budget (fair-device ops are
1626/// ≤~100 ms even at 1024px) means another process owns the device —
1627/// verdicts are per-process, so CPU for the rest of this one.
1628static MM_KILL: AtomicBool = AtomicBool::new(false);
1629pub(crate) fn mm_killed() -> bool {
1630    MM_KILL.load(Ordering::Relaxed)
1631}
1632pub(crate) fn mm_kill() {
1633    MM_KILL.store(true, Ordering::Relaxed);
1634}
1635
1636/// Consecutive over-budget ops. ONE slow op is not contention: on a
1637/// 24 GB Mac running the 25.7 GB fl2va file the first ops after the
1638/// prompt encode page their weights in from the SSD and take seconds —
1639/// a field report (hololabs, HF discussion #2) had to neuter the kill
1640/// to keep the denoise on the GPU, and then measured 48 s/step where the
1641/// CPU fallback took >60. Contention is persistent; a page-in is not.
1642static MM_STRIKES: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
1643const MM_STRIKES_TO_KILL: u32 = 3;
1644/// Whether the kill is armed at all. A one-shot phase whose slowness is
1645/// expected and not contention — the video prompt encoder streaming
1646/// 12 GB off the SSD on a 24 GB Mac (HF discussion #4: users had to
1647/// gut `mm_kill` to keep the denoise loop on the GPU) — disarms it and
1648/// re-arms it when the phase is over; strikes taken meanwhile are
1649/// forgotten.
1650static MM_ARMED: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
1651
1652/// Disarm / re-arm the contention kill around a phase whose GEMMs are
1653/// slow for reasons that are not another process (see `MM_ARMED`).
1654pub fn mm_kill_arm(on: bool) {
1655    MM_ARMED.store(on, Ordering::Relaxed);
1656    if on {
1657        MM_STRIKES.store(0, Ordering::Relaxed);
1658    }
1659}
1660
1661/// The contention verdict for one wide op: `el` against its
1662/// work-proportional `budget`. `exempt` marks ops whose time is not
1663/// evidence — the cold probe, or a weight that was not resident before
1664/// the call and rode in with it. Kills after `MM_STRIKES_TO_KILL`
1665/// consecutive strikes; a within-budget op clears the count.
1666/// `CMF_MM_KILL=0` disables the kill entirely (the device is trusted).
1667pub(crate) fn mm_budget_check(
1668    what: &str,
1669    el: std::time::Duration,
1670    budget: std::time::Duration,
1671    exempt: bool,
1672) {
1673    if el <= budget {
1674        MM_STRIKES.store(0, Ordering::Relaxed);
1675        return;
1676    }
1677    if exempt || !MM_ARMED.load(Ordering::Relaxed) {
1678        return;
1679    }
1680    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1681    let on = *ON.get_or_init(|| std::env::var("CMF_MM_KILL").as_deref() != Ok("0"));
1682    let n = MM_STRIKES.fetch_add(1, Ordering::Relaxed) + 1;
1683    if !on {
1684        tracing::info!(
1685            "gpu {what} took {el:?} (budget {budget:?}) — over budget, CMF_MM_KILL=0 keeps the device"
1686        );
1687        return;
1688    }
1689    if n >= MM_STRIKES_TO_KILL {
1690        tracing::warn!(
1691            "gpu {what} took {el:?} (budget {budget:?}), {n} in a row — \
1692             device contended, CPU for the rest of the process (CMF_MM_KILL=0 to override)"
1693        );
1694        mm_kill();
1695    } else {
1696        tracing::info!(
1697            "gpu {what} took {el:?} (budget {budget:?}) — strike {n} of {MM_STRIKES_TO_KILL}"
1698        );
1699    }
1700}
1701
1702/// Fused DiT SwiGLU FFN on the device: g=X·W1ᵀ, u=X·W3ᵀ, silu(g)·u,
1703/// Causal chunk attention on the device: `b` queries against `s0 + b`
1704/// cached keys. wgpu only — Metal's chunk graph keeps attention inside
1705/// the resident block and never calls out.
1706#[allow(unused_variables, clippy::too_many_arguments)]
1707pub fn chunk_attend(
1708    q: &[f32],
1709    k: &[&[f32]],
1710    v: &[&[f32]],
1711    b: usize,
1712    s0: usize,
1713    nh: usize,
1714    nkv: usize,
1715    hd: usize,
1716    scale: f32,
1717    out: &mut [f32],
1718) -> bool {
1719    match backend() {
1720        #[cfg(feature = "gpu")]
1721        Backend::Wgpu => crate::gpu_wgpu::chunk_attend(q, k, v, b, s0, nh, nkv, hd, scale, out),
1722        #[allow(unreachable_patterns)]
1723        _ => false,
1724    }
1725}
1726
1727/// Fused QKV projection: one upload of the normed chunk, three GEMMs,
1728/// one readback of Q|K|V back to back. Metal has no twin yet — its
1729/// chunk graph keeps the whole layer resident and never surfaces QKV.
1730#[allow(unused_variables, clippy::too_many_arguments)]
1731pub fn q4t_qkv(
1732    model: &Arc<CmfModel>,
1733    wq: usize,
1734    wk: usize,
1735    wv: usize,
1736    xs: &[f32],
1737    b: usize,
1738    cols: usize,
1739    rq: usize,
1740    rk: usize,
1741    rv: usize,
1742    out: &mut [f32],
1743) -> bool {
1744    match backend() {
1745        #[cfg(feature = "gpu")]
1746        Backend::Wgpu => crate::gpu_wgpu::q4t_qkv(model, wq, wk, wv, xs, b, cols, rq, rk, rv, out),
1747        #[allow(unreachable_patterns)]
1748        _ => false,
1749    }
1750}
1751
1752/// y=·W2ᵀ — one command buffer, only X and Y cross the CPU boundary.
1753#[allow(unused_variables, clippy::too_many_arguments)]
1754/// SwiGLU FFN with a row-packed [gate|up] fc1 (MiniMax-H3's DiT), run
1755/// end to end on the device. wgpu only: Metal keeps the host loop until
1756/// its own packed kernel exists.
1757#[allow(clippy::too_many_arguments, unused_variables)]
1758pub fn q4tp_ffn_packed(
1759    model: &Arc<CmfModel>,
1760    w1: usize,
1761    w2: usize,
1762    xs: &[f32],
1763    b: usize,
1764    hidden: usize,
1765    inter: usize,
1766    bias: Option<&[f32]>,
1767    out: &mut [f32],
1768) -> bool {
1769    match backend() {
1770        #[cfg(feature = "gpu")]
1771        Backend::Wgpu => {
1772            crate::gpu_wgpu::ffn_packed(model, w1, w2, xs, b, hidden, inter, bias, out)
1773        }
1774        #[allow(unreachable_patterns)]
1775        _ => false,
1776    }
1777}
1778
1779pub fn q4tp_ffn(
1780    model: &Arc<CmfModel>,
1781    w1: usize,
1782    w3: usize,
1783    w2: usize,
1784    xs: &[f32],
1785    b: usize,
1786    hidden: usize,
1787    inter: usize,
1788    out: &mut [f32],
1789) -> bool {
1790    match backend() {
1791        #[cfg(target_os = "macos")]
1792        Backend::Metal => crate::gpu_metal::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
1793        #[cfg(feature = "gpu")]
1794        Backend::Wgpu => crate::gpu_wgpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
1795        #[allow(unreachable_patterns)]
1796        _ => false,
1797    }
1798}
1799
1800pub fn q4t_ffn(
1801    model: &Arc<CmfModel>,
1802    w1: usize,
1803    w3: usize,
1804    w2: usize,
1805    xs: &[f32],
1806    b: usize,
1807    hidden: usize,
1808    inter: usize,
1809    out: &mut [f32],
1810) -> bool {
1811    match backend() {
1812        #[cfg(target_os = "macos")]
1813        Backend::Metal => crate::gpu_metal::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
1814        #[cfg(feature = "gpu")]
1815        Backend::Wgpu => crate::gpu_wgpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
1816        #[allow(unreachable_patterns)]
1817        _ => false,
1818    }
1819}
1820
1821/// One whole modulated DiT block for `dit_block`: geometry, norm
1822/// weights, AdaLN scale/gate vectors (gates pre-tanh'd), a per-token
1823/// f32 RoPE cos/sin table, and the directory indices of the seven
1824/// q4t projections. `x` is in-out `[n, hidden]`.
1825pub struct DitBlockArgs<'a> {
1826    pub n: usize,
1827    pub hidden: usize,
1828    pub inter: usize,
1829    pub nh: usize,
1830    pub nkv: usize,
1831    pub hd: usize,
1832    pub eps: f32,
1833    pub rope_cos: &'a [f32],
1834    pub rope_sin: &'a [f32],
1835    pub norm1: &'a [f32],
1836    pub norm2: &'a [f32],
1837    pub ffn_norm1: &'a [f32],
1838    pub ffn_norm2: &'a [f32],
1839    pub norm_q: &'a [f32],
1840    pub norm_k: &'a [f32],
1841    pub s_msa: &'a [f32],
1842    pub gate_msa: &'a [f32],
1843    pub s_mlp: &'a [f32],
1844    pub gate_mlp: &'a [f32],
1845    pub wq: usize,
1846    pub wk: usize,
1847    pub wv: usize,
1848    pub wo: usize,
1849    pub w1: usize,
1850    pub w3: usize,
1851    pub w2: usize,
1852    /// The projections' layout: q4tp (ladder scales) vs plain q4_tiled.
1853    /// The recommended Lumina file is q4tp, and a backend that only
1854    /// knows q4t must decline rather than decode with the wrong reader.
1855    pub q4tp: bool,
1856    /// The hidden state is already on the device from the previous block,
1857    /// so `x` need not be uploaded.
1858    pub resident_in: bool,
1859    /// Leave the result on the device instead of reading it back. The DiT
1860    /// loop does not touch `x` between blocks, so 27 of every 28 readbacks
1861    /// were moving 19 MB across PCIe and stalling on it for nothing.
1862    pub resident_out: bool,
1863}
1864
1865/// Can the selected backend keep the DiT's hidden state on the device
1866/// between blocks? Only the wgpu whole-block path; the Metal entry takes
1867/// and returns host memory every call.
1868pub fn dit_chain_supported() -> bool {
1869    #[cfg(feature = "gpu")]
1870    {
1871        return matches!(backend(), Backend::Wgpu) && fused_dit_block_available();
1872    }
1873    #[allow(unreachable_code)]
1874    false
1875}
1876
1877/// Pull the resident hidden state back to the host. For the caller that
1878/// chained blocks and then hit one the device declined.
1879pub fn dit_state_fetch(_x: &mut [f32]) -> bool {
1880    #[cfg(feature = "gpu")]
1881    {
1882        if matches!(backend(), Backend::Wgpu) {
1883            return crate::gpu_wgpu::dit_state_fetch(_x);
1884        }
1885    }
1886    false
1887}
1888
1889/// One whole modulated DiT block on the device — norms, qkv, RoPE,
1890/// attention, residuals and the SwiGLU FFN in a single command
1891/// buffer; only `x` crosses the CPU boundary (in and out).
1892#[allow(unused_variables)]
1893/// The DiT's three projections in one submission (wgpu only; the
1894/// Metal path fuses the whole block instead). False = the caller keeps
1895/// its three separate calls.
1896#[allow(unused_variables, clippy::too_many_arguments)]
1897pub fn dit_qkv(
1898    model: &Arc<CmfModel>,
1899    wq: usize,
1900    wk: usize,
1901    wv: usize,
1902    xs: &[f32],
1903    b: usize,
1904    hidden: usize,
1905    qrows: usize,
1906    kvrows: usize,
1907    q_out: &mut [f32],
1908    k_out: &mut [f32],
1909    v_out: &mut [f32],
1910) -> bool {
1911    match backend() {
1912        #[cfg(feature = "gpu")]
1913        Backend::Wgpu => crate::gpu_wgpu::q4tp_qkv(
1914            model, wq, wk, wv, xs, b, hidden, qrows, kvrows, q_out, k_out, v_out,
1915        ),
1916        #[allow(unreachable_patterns)]
1917        _ => false,
1918    }
1919}
1920
1921/// Is a FUSED whole-block device path on offer? The batched-CFG shape
1922/// (two sequences in one tall batch) and the fused block (one sequence,
1923/// one command buffer) are alternatives, and the caller picks.
1924pub fn fused_dit_block_available() -> bool {
1925    #[cfg(target_os = "macos")]
1926    {
1927        matches!(backend(), Backend::Metal) && fused_block_trusted()
1928    }
1929    #[cfg(not(target_os = "macos"))]
1930    {
1931        false
1932    }
1933}
1934
1935pub fn dit_block(model: &Arc<CmfModel>, a: &DitBlockArgs, x: &mut [f32]) -> bool {
1936    dit_block_seg(model, a, &[a.n], x)
1937}
1938
1939/// The same block over a CONCATENATION of independent sequences:
1940/// attention per segment, everything position-wise batched. wgpu only —
1941/// the Metal path takes the single-sequence entry above.
1942pub fn dit_block_seg(
1943    model: &Arc<CmfModel>,
1944    a: &DitBlockArgs,
1945    segs: &[usize],
1946    x: &mut [f32],
1947) -> bool {
1948    match backend() {
1949        #[cfg(target_os = "macos")]
1950        Backend::Metal if segs.len() <= 1 => crate::gpu_metal::dit_block(model, a, x),
1951        // The wgpu whole-block path. What it buys is host round trips —
1952        // six a block become one — so it defaults ON where those cost
1953        // real time (a discrete card across PCIe) and OFF on unified
1954        // memory, where the per-op path shares the same pages and the
1955        // fusion measured slightly slower on an M4. `CMF_DIT_FUSED=1`
1956        // forces it anywhere, `=0` forbids it.
1957        #[cfg(feature = "gpu")]
1958        Backend::Wgpu
1959            if match std::env::var("CMF_DIT_FUSED").ok().as_deref() {
1960                Some("0") => false,
1961                Some(_) => true,
1962                None => crate::gpu_wgpu::discrete_active(),
1963            } =>
1964        {
1965            crate::gpu_wgpu::dit_block_seg(model, a, segs, x)
1966        }
1967        #[allow(unreachable_patterns)]
1968        _ => false,
1969    }
1970}
1971
1972/// One VAE resnet block for `vae_resnet`: norm/conv weights and the
1973/// channel/shape geometry. `shortcut` is the 1×1 projection (w, b, k)
1974/// when in/out channels differ.
1975pub struct VaeResnetArgs<'a> {
1976    pub groups: usize,
1977    pub ic: usize,
1978    pub oc: usize,
1979    pub h: usize,
1980    pub w: usize,
1981    pub n1w: &'a [f32],
1982    pub n1b: &'a [f32],
1983    pub c1w: &'a [f32],
1984    pub c1b: &'a [f32],
1985    pub c1k: usize,
1986    pub n2w: &'a [f32],
1987    pub n2b: &'a [f32],
1988    pub c2w: &'a [f32],
1989    pub c2b: &'a [f32],
1990    pub c2k: usize,
1991    pub shortcut: Option<(&'a [f32], &'a [f32], usize)>,
1992}
1993
1994/// One whole VAE resnet block on the device (norm+silu → conv ×2 →
1995/// shortcut → add, one command buffer).
1996#[allow(unused_variables)]
1997pub fn vae_resnet(a: &VaeResnetArgs, x: &[f32], out: &mut [f32]) -> bool {
1998    match backend() {
1999        #[cfg(target_os = "macos")]
2000        Backend::Metal => crate::gpu_metal::vae_resnet(a, x, out),
2001        _ => false,
2002    }
2003}
2004
2005/// Nearest-2× upsample fused with the following conv — the small
2006/// pre-upsample image is what crosses the CPU boundary.
2007#[allow(unused_variables, clippy::too_many_arguments)]
2008pub fn vae_upsample_conv(
2009    w: &[f32],
2010    bias: &[f32],
2011    x: &[f32],
2012    ic: usize,
2013    oc: usize,
2014    h: usize,
2015    w_img: usize,
2016    k: usize,
2017    out: &mut [f32],
2018) -> bool {
2019    match backend() {
2020        #[cfg(target_os = "macos")]
2021        Backend::Metal => crate::gpu_metal::vae_upsample_conv(w, bias, x, ic, oc, h, w_img, k, out),
2022        #[cfg(feature = "gpu")]
2023        Backend::Wgpu => crate::gpu_wgpu::vae_upsample_conv(w, bias, x, ic, oc, h, w_img, k, out),
2024        #[allow(unreachable_patterns)]
2025        _ => false,
2026    }
2027}
2028
2029/// VAE conv2d on the device (implicit GEMM — the CPU path pays for a
2030/// multi-GB im2col matrix at high resolutions).
2031#[allow(unused_variables, clippy::too_many_arguments)]
2032pub fn vae_conv2d(
2033    w: &[f32],
2034    bias: &[f32],
2035    x: &[f32],
2036    ic: usize,
2037    oc: usize,
2038    h: usize,
2039    w_img: usize,
2040    k: usize,
2041    out: &mut [f32],
2042) -> bool {
2043    match backend() {
2044        #[cfg(target_os = "macos")]
2045        Backend::Metal => crate::gpu_metal::vae_conv2d(w, bias, x, ic, oc, h, w_img, k, out),
2046        #[cfg(feature = "gpu")]
2047        Backend::Wgpu => crate::gpu_wgpu::vae_conv2d(w, bias, x, ic, oc, h, w_img, k, out),
2048        #[allow(unreachable_patterns)]
2049        _ => false,
2050    }
2051}
2052
2053/// DiT full bidirectional attention on the device (all heads:
2054/// scores GEMM → row softmax → P·V → panel unstack, one command
2055/// buffer). Head-major inputs; out is [n, nh·hd].
2056#[allow(unused_variables, clippy::too_many_arguments)]
2057/// Attention from an interleaved qkv panel, splitting into head-major
2058/// planes ON the device. wgpu only; `false` elsewhere so the caller
2059/// keeps its host repack.
2060#[allow(unused_variables)]
2061#[allow(clippy::too_many_arguments)]
2062/// qkv projection + attention with the panel never leaving the card.
2063/// wgpu only; `false` elsewhere and the caller keeps its host chain.
2064#[allow(clippy::too_many_arguments, unused_variables)]
2065pub fn dit_qkv_attention(
2066    model: &Arc<CmfModel>,
2067    qkv_idx: usize,
2068    xn: &[f32],
2069    n: usize,
2070    hidden: usize,
2071    nh: usize,
2072    hd: usize,
2073    scale: f32,
2074    nr: (&[f32], &[f32], &[f32], f32),
2075    out: &mut [f32],
2076) -> bool {
2077    match backend() {
2078        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2079        Backend::Wgpu => crate::gpu_wgpu::dit_qkv_attention(
2080            model, qkv_idx, xn, n, hidden, nh, hd, scale, nr, out,
2081        ),
2082        #[allow(unreachable_patterns)]
2083        _ => false,
2084    }
2085}
2086
2087/// The whole attention half of a DiT block on the card: qkv GEMM,
2088/// attention, output projection. Only `proj` comes home.
2089#[allow(clippy::too_many_arguments)]
2090pub fn dit_qkv_attn_out(
2091    model: &Arc<CmfModel>,
2092    qkv_idx: usize,
2093    out_idx: usize,
2094    xn: &[f32],
2095    n: usize,
2096    hidden: usize,
2097    nh: usize,
2098    hd: usize,
2099    scale: f32,
2100    nr: (&[f32], &[f32], &[f32], f32),
2101    proj: &mut [f32],
2102) -> bool {
2103    match backend() {
2104        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2105        Backend::Wgpu => crate::gpu_wgpu::dit_qkv_attn_out(
2106            model, qkv_idx, out_idx, xn, n, hidden, nh, hd, scale, nr, proj,
2107        ),
2108        #[allow(unreachable_patterns)]
2109        _ => false,
2110    }
2111}
2112
2113/// The VAE decoder's attention half on the card. Only `proj` returns.
2114#[allow(clippy::too_many_arguments)]
2115pub fn vae_qkv_attn_out(
2116    model: &Arc<CmfModel>,
2117    qkv_idx: usize,
2118    out_idx: usize,
2119    xn: &[f32],
2120    n: usize,
2121    dim: usize,
2122    nh: usize,
2123    hd: usize,
2124    scale: f32,
2125    angles: &[f32],
2126    eps: f32,
2127    qkv_bias: &[f32],
2128    proj: &mut [f32],
2129) -> bool {
2130    match backend() {
2131        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2132        Backend::Wgpu => crate::gpu_wgpu::vae_qkv_attn_out(
2133            model, qkv_idx, out_idx, xn, n, dim, nh, hd, scale, angles, eps, qkv_bias, proj,
2134        ),
2135        #[allow(unreachable_patterns)]
2136        _ => false,
2137    }
2138}
2139
2140#[allow(clippy::too_many_arguments)]
2141pub fn vae_attention_packed(
2142    qkv: &[f32],
2143    nh: usize,
2144    n: usize,
2145    hd: usize,
2146    scale: f32,
2147    angles: &[f32],
2148    eps: f32,
2149    out: &mut [f32],
2150) -> bool {
2151    vae_attention_packed_layout(qkv, nh, n, hd, scale, angles, eps, out, 1)
2152}
2153
2154#[allow(clippy::too_many_arguments)]
2155pub fn vae_attention_packed_layout(
2156    qkv: &[f32],
2157    nh: usize,
2158    n: usize,
2159    hd: usize,
2160    scale: f32,
2161    angles: &[f32],
2162    eps: f32,
2163    out: &mut [f32],
2164    layout: u32,
2165) -> bool {
2166    match backend() {
2167        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2168        Backend::Wgpu => crate::gpu_wgpu::vae_attention_packed_layout(
2169            qkv, nh, n, hd, scale, angles, eps, out, layout,
2170        ),
2171        #[allow(unreachable_patterns)]
2172        _ => false,
2173    }
2174}
2175
2176#[allow(clippy::too_many_arguments)]
2177pub fn dit_split_only(
2178    qkv: &[f32],
2179    nh: usize,
2180    n: usize,
2181    hd: usize,
2182    layout: u32,
2183    norm: Option<(&[f32], f32)>,
2184    out_q: &mut [f32],
2185) -> bool {
2186    match backend() {
2187        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2188        Backend::Wgpu => crate::gpu_wgpu::dit_split_only(qkv, nh, n, hd, layout, norm, out_q),
2189        #[allow(unreachable_patterns)]
2190        _ => false,
2191    }
2192}
2193
2194/// The backend's f32 NT GEMM: `y[n×m] = x[n×k] · wᵀ[m×k]`. Tensor
2195/// cores where the card has them. Refuses under `CMF_BAKE_GPU=0` or
2196/// strict f32, and for jobs below n·k·m = 4M, where the round trip
2197/// costs more than the arithmetic saves.
2198/// `gemm_nt_f32` whose `w` is known to change every call (an
2199/// accumulation over fresh activations, not a weight): it skips the
2200/// resident ledger and its per-call fingerprint of the whole operand.
2201pub fn gemm_nt_f32_transient(
2202    x: &[f32],
2203    w: &[f32],
2204    y: &mut [f32],
2205    n: usize,
2206    k: usize,
2207    m: usize,
2208) -> bool {
2209    match backend() {
2210        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2211        Backend::Wgpu => crate::gpu_wgpu::gemm_nt_f32_transient(x, w, y, n, k, m),
2212        #[allow(unreachable_patterns)]
2213        _ => false,
2214    }
2215}
2216
2217pub fn gemm_nt_f32(x: &[f32], w: &[f32], y: &mut [f32], n: usize, k: usize, m: usize) -> bool {
2218    match backend() {
2219        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2220        Backend::Wgpu => crate::gpu_wgpu::gemm_nt_f32(x, w, y, n, k, m),
2221        #[allow(unreachable_patterns)]
2222        _ => false,
2223    }
2224}
2225
2226/// Music-3's FFN chain resident on the device — two GEMMs and the GLU
2227/// between them with no host round trip. `false` = refused, host runs.
2228#[allow(clippy::too_many_arguments)]
2229pub fn music3_ffn(
2230    model: &std::sync::Arc<CmfModel>,
2231    idx_in: usize,
2232    idx_out: usize,
2233    h: &[f32],
2234    bias_in: &[f32],
2235    n: usize,
2236    hs: usize,
2237    inter: usize,
2238    out: &mut [f32],
2239) -> bool {
2240    match backend() {
2241        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2242        Backend::Wgpu => {
2243            crate::gpu_wgpu::music3_ffn(model, idx_in, idx_out, h, bias_in, n, hs, inter, out)
2244        }
2245        #[allow(unreachable_patterns)]
2246        _ => false,
2247    }
2248}
2249
2250/// A 1D convolution as a GEMM whose column matrix is expanded on the
2251/// device instead of being built, transposed and uploaded by the host.
2252/// `yt` comes back `[out_n x oc]`. `false` = refused, caller runs host.
2253#[allow(clippy::too_many_arguments)]
2254pub fn conv1d_gemm(
2255    x: &[f32],
2256    w: &[f32],
2257    ic: usize,
2258    oc: usize,
2259    n: usize,
2260    k: usize,
2261    pad: usize,
2262    dil: usize,
2263    out_n: usize,
2264    yt: &mut [f32],
2265) -> bool {
2266    match backend() {
2267        #[cfg(target_os = "macos")]
2268        Backend::Metal => crate::gpu_metal::conv1d_gemm(x, w, ic, oc, n, k, pad, dil, out_n, yt),
2269        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2270        Backend::Wgpu => crate::gpu_wgpu::conv1d_gemm(x, w, ic, oc, n, k, pad, dil, out_n, yt),
2271        #[allow(unreachable_patterns)]
2272        _ => false,
2273    }
2274}
2275
2276/// The convolution as a GEMM on the matrix units. `false` = refused.
2277#[allow(clippy::too_many_arguments)]
2278pub fn vae_conv2d_coop(
2279    w: &[f32],
2280    bias: Option<&[f32]>,
2281    x: &[f32],
2282    ic: usize,
2283    oc: usize,
2284    h: usize,
2285    wi: usize,
2286    k: usize,
2287    out: &mut [f32],
2288) -> bool {
2289    match backend() {
2290        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2291        Backend::Wgpu => crate::gpu_wgpu::vae_conv2d_coop(w, bias, x, ic, oc, h, wi, k, out),
2292        #[allow(unreachable_patterns)]
2293        _ => false,
2294    }
2295}
2296
2297pub fn dit_attention_packed(
2298    qkv: &[f32],
2299    nh: usize,
2300    n: usize,
2301    hd: usize,
2302    scale: f32,
2303    // (rope angles, q norm weights, k norm weights, eps) when the device
2304    // should apply qk-norm and RoPE itself; None when the host already did.
2305    nr: Option<(&[f32], &[f32], &[f32], f32)>,
2306    out: &mut [f32],
2307) -> bool {
2308    match backend() {
2309        // wgpu carries the only implementation, and it is not
2310        // platform-specific: `CMF_GPU=wgpu` on macOS runs it over Metal
2311        // like anywhere else. It used to be compiled out here on macOS,
2312        // which made the call a silent `false` — and the caller's
2313        // `assert!` turned that refusal into a panic on every
2314        // `cortiq animate` this platform ever ran.
2315        #[cfg(feature = "gpu")]
2316        Backend::Wgpu => crate::gpu_wgpu::dit_attention_packed(qkv, nh, n, hd, scale, nr, out),
2317        #[allow(unreachable_patterns)]
2318        _ => false,
2319    }
2320}
2321
2322/// Whether `dit_attention_packed` has an implementation on the backend
2323/// that is actually selected.
2324///
2325/// The caller has to know BEFORE it skips the host qk-norm: deferring
2326/// the norm to a device that then refuses leaves q/k unnormalized with
2327/// no way back. Native Metal has no packed kernel, so on macOS this is
2328/// false unless `CMF_GPU=wgpu` picked the other backend.
2329pub fn dit_attention_packed_available() -> bool {
2330    #[allow(unreachable_patterns)]
2331    match backend() {
2332        #[cfg(feature = "gpu")]
2333        Backend::Wgpu => crate::gpu_wgpu::dit_attention_packed_ready(),
2334        _ => false,
2335    }
2336}
2337
2338pub fn dit_attention(
2339    qh: &[f32],
2340    kh: &[f32],
2341    vh: &[f32],
2342    nh: usize,
2343    nkv: usize,
2344    n: usize,
2345    hd: usize,
2346    scale: f32,
2347    out: &mut [f32],
2348) -> bool {
2349    match backend() {
2350        #[cfg(target_os = "macos")]
2351        Backend::Metal => crate::gpu_metal::dit_attention(qh, kh, vh, nh, nkv, n, hd, scale, out),
2352        #[cfg(feature = "gpu")]
2353        Backend::Wgpu => crate::gpu_wgpu::dit_attention(qh, kh, vh, nh, nkv, n, hd, scale, out),
2354        #[allow(unreachable_patterns)]
2355        _ => false,
2356    }
2357}
2358
2359/// Batched q4t GEMM on the device (imagegen DiT prefill shapes).
2360/// Metal: q4t_mul_mm decodes the mmap-resident tiles inside the
2361/// GEMM's K loop. wgpu (Vulkan/DX12 → NVIDIA/AMD/Intel/Adreno/Mali):
2362/// the register-blocked WGSL twin, weights cached in VRAM.
2363#[allow(unused_variables)]
2364pub fn q4tp_matmat(
2365    model: &Arc<CmfModel>,
2366    idx: usize,
2367    xs: &[f32],
2368    b: usize,
2369    rows: usize,
2370    cols: usize,
2371    out: &mut [f32],
2372) -> bool {
2373    match backend() {
2374        #[cfg(target_os = "macos")]
2375        Backend::Metal => crate::gpu_metal::q4tp_matmat(model, idx, xs, b, rows, cols, out),
2376        #[cfg(feature = "gpu")]
2377        Backend::Wgpu => crate::gpu_wgpu::q4tp_matmat(model, idx, xs, b, rows, cols, out),
2378        #[allow(unreachable_patterns)]
2379        _ => false,
2380    }
2381}
2382
2383/// The same over a two-bit weight plane. Metal has no q2tp kernel, so
2384/// there it declines and the host takes it.
2385pub fn q2tp_matmat(
2386    model: &Arc<CmfModel>,
2387    idx: usize,
2388    xs: &[f32],
2389    b: usize,
2390    rows: usize,
2391    cols: usize,
2392    out: &mut [f32],
2393) -> bool {
2394    match backend() {
2395        #[cfg(feature = "gpu")]
2396        Backend::Wgpu => crate::gpu_wgpu::q2tp_matmat(model, idx, xs, b, rows, cols, out),
2397        #[allow(unreachable_patterns)]
2398        _ => false,
2399    }
2400}
2401
2402/// Single-token q4tp matvec on the device — the lm_head class. Through the
2403/// DEDICATED matvec kernel: the batched GEMM at b=1 measured 11.73 ms
2404/// against the host's 9.51 on the release head, so the route that was
2405/// supposed to save eleven milliseconds a token lost its own probe instead.
2406pub fn q4tp_matvec(
2407    model: &Arc<CmfModel>,
2408    idx: usize,
2409    xs: &[f32],
2410    rows: usize,
2411    cols: usize,
2412    out: &mut [f32],
2413) -> bool {
2414    match backend() {
2415        #[cfg(target_os = "macos")]
2416        Backend::Metal => crate::gpu_metal::q4tp_matvec_for_test(model, idx, xs, rows, cols, out),
2417        #[cfg(feature = "gpu")]
2418        Backend::Wgpu => crate::gpu_wgpu::q4tp_matvec(model, idx, xs, rows, cols, out),
2419        #[allow(unreachable_patterns)]
2420        _ => false,
2421    }
2422}
2423
2424/// Single-token q4_tiled matvec on the device — the lm_head class (a
2425/// q4t checkpoint's head is its biggest host matvec, exactly like the
2426/// q4tp twin above). wgpu holds q4t_mv pipelines only inside the graph
2427/// encoder — the standalone arm stays an honest refusal until a
2428/// discrete-GPU q4t model reaches the bench.
2429pub fn q4t_matvec(
2430    model: &Arc<CmfModel>,
2431    idx: usize,
2432    xs: &[f32],
2433    rows: usize,
2434    cols: usize,
2435    out: &mut [f32],
2436) -> bool {
2437    match backend() {
2438        #[cfg(target_os = "macos")]
2439        Backend::Metal => crate::gpu_metal::q4t_matvec_for_test(model, idx, xs, rows, cols, out),
2440        #[allow(unreachable_patterns)]
2441        _ => false,
2442    }
2443}
2444
2445pub fn q4t_matmat(
2446    model: &Arc<CmfModel>,
2447    idx: usize,
2448    xs: &[f32],
2449    b: usize,
2450    rows: usize,
2451    cols: usize,
2452    out: &mut [f32],
2453) -> bool {
2454    match backend() {
2455        #[cfg(target_os = "macos")]
2456        Backend::Metal => crate::gpu_metal::q4t_matmat(model, idx, xs, b, rows, cols, out),
2457        #[cfg(feature = "gpu")]
2458        Backend::Wgpu => crate::gpu_wgpu::q4t_matmat(model, idx, xs, b, rows, cols, out),
2459        #[allow(unreachable_patterns)]
2460        _ => false,
2461    }
2462}
2463
2464/// Whole-block token-graph types re-exported from the Metal backend.
2465#[cfg(target_os = "macos")]
2466pub use crate::gpu_metal::{
2467    AttnDeviceParams, AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GpuMoe, GraphDims, MetalFfn,
2468    O1AttnParams, TokenGraph, kv_mirror_drop, kv_mirror_read_last, kv_mirror_take_imp,
2469};
2470
2471/// A BLOCK of consecutive q1 GDN layers in one submission (Metal only).
2472#[cfg(target_os = "macos")]
2473pub fn gdn_block(
2474    model: &Arc<CmfModel>,
2475    layers: &[GdnGpuLayer],
2476    states: &mut [&mut [f32]],
2477    cfg: &GdnGpuCfg,
2478    h: &mut [f32],
2479) -> bool {
2480    match backend() {
2481        Backend::Metal => crate::gpu_metal::gdn_block(model, layers, states, cfg, h),
2482        _ => false,
2483    }
2484}
2485
2486/// A layer's MoE-FFN in one submission (amortizing the dispatch cost).
2487#[allow(unused_variables)]
2488pub fn moe_block(model: &Arc<CmfModel>, jobs: &[MoeJob], out: &mut [f32]) -> bool {
2489    match backend() {
2490        #[cfg(target_os = "macos")]
2491        Backend::Metal => crate::gpu_metal::moe_block(model, jobs, out),
2492        #[cfg(feature = "gpu")]
2493        Backend::Wgpu => crate::gpu_wgpu::moe_block(model, jobs, out),
2494        Backend::None => false,
2495    }
2496}
2497
2498/// Independent matvecs of one input in a single submission (GDN projections).
2499#[allow(unused_variables)]
2500pub fn matvec_batch(model: &Arc<CmfModel>, jobs: &[BatchJob], out: &mut [&mut [f32]]) -> bool {
2501    match backend() {
2502        #[cfg(target_os = "macos")]
2503        Backend::Metal => crate::gpu_metal::matvec_batch(model, jobs, out),
2504        #[cfg(feature = "gpu")]
2505        Backend::Wgpu => crate::gpu_wgpu::matvec_batch(model, jobs, out),
2506        Backend::None => false,
2507    }
2508}
2509
2510// ── Whole-token wgpu graph race (generation granularity) ─────────────
2511// On integrated/mobile adapters the graph is neither trusted nor banned
2512// a priori — it RACES the normal path: generations alternate arms (the
2513// normal path first — known-good UX — then the graph), per-token wall
2514// times accumulate per arm, and once both arms have enough steady
2515// samples the faster one wins for the process. Arm switches happen ONLY
2516// at generation boundaries (`kv_cache.clear()` resets state), so the
2517// device KV mirror and the CPU cache never diverge mid-sequence. The
2518// single exception is the first-token bail: the very first decode token
2519// of a graph generation may be discarded and recomputed on the CPU
2520// path (the prompt KV is CPU-owned at that point, so this is safe) —
2521// a tiled mobile GPU that drains its pipeline at every barrier turns
2522// the ~300-dispatch graph into seconds per token (field report: 0.2
2523// tok/s vs 15 on the CPU), and one token is all it takes to see that.
2524static GRAPH_RACE_STATE: AtomicU8 = AtomicU8::new(0); // 0 racing, 1 graph won, 2 normal won
2525static GRAPH_RACE_FLIP: AtomicU32 = AtomicU32::new(0);
2526static GRAPH_RACE_ARM_GRAPH: AtomicU8 = AtomicU8::new(0); // this generation's arm
2527static GRAPH_RACE_TOK: AtomicU32 = AtomicU32::new(0); // token index within the generation
2528static GRAPH_NS: [AtomicU64; 2] = [AtomicU64::new(0), AtomicU64::new(0)]; // [normal, graph]
2529static GRAPH_N: [AtomicU32; 2] = [AtomicU32::new(0), AtomicU32::new(0)];
2530
2531/// Steady per-token samples per arm before the race decides.
2532const GRAPH_RACE_SAMPLES: u32 = 4;
2533
2534/// Called at every generation start (fresh KV). Applies a pending
2535/// verdict and picks this generation's arm while racing.
2536/// A graph that cannot be built for THIS model will never build: the
2537/// refusal is a property of the weights, not of the moment. Retrying it
2538/// per token is not free — the builder walks every layer and asks each
2539/// tensor for a graph view before giving up at layer 0 — and on an
2540/// Adreno 642L that retry cost 3x: forcing the graph on a model it
2541/// refuses measured 0.3 tok/s against 0.905 for the per-op path it falls
2542/// back to. Remembered once, the fallback runs at its own speed.
2543static GRAPH_UNSUPPORTED: AtomicBool = AtomicBool::new(false);
2544
2545/// The builder refused for a STRUCTURAL reason — an unsupported weight
2546/// or layer kind. Callers must NOT report the transient refusals (an
2547/// unsealed o1 state during prefill, a softcap): those clear on their
2548/// own and marking them would disable the graph for good.
2549pub fn graph_mark_unsupported() {
2550    if !GRAPH_UNSUPPORTED.swap(true, Ordering::Relaxed) {
2551        tracing::info!("wgpu token graph: unsupported for this model — not retrying");
2552    }
2553}
2554
2555pub fn graph_unsupported() -> bool {
2556    GRAPH_UNSUPPORTED.load(Ordering::Relaxed)
2557}
2558
2559/// A different model in the same process starts with a clean slate.
2560pub fn graph_unsupported_reset() {
2561    GRAPH_UNSUPPORTED.store(false, Ordering::Relaxed);
2562}
2563
2564pub fn graph_race_begin_generation() {
2565    // One generation has now compiled whatever this model needs; keep it
2566    // for the next process. Once per run: the blob does not grow after
2567    // the pipelines exist, and the write is megabytes against the ~200 s
2568    // of compiling it saves on the device that needed this.
2569    #[cfg(feature = "gpu")]
2570    {
2571        // Save once, at the start of the SECOND generation: the first
2572        // has dispatched, so there is something to keep, and nothing is
2573        // saved before any work (the driver compiles at first use, not
2574        // at pipeline creation — the context comes up in 1.5 s while the
2575        // compiling costs minutes).
2576        //
2577        // Flushing again on 4, 8, 16 … was tried on the theory that a
2578        // chat turn compiles shapes the first one did not. It buys
2579        // nothing: a fresh app process still spent 49.0 s, then 58.7,
2580        // then 61.3 on its first answer with the backoff in place. One
2581        // flush it is.
2582        static FLUSHED: std::sync::Once = std::sync::Once::new();
2583        static FIRST: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
2584        if FIRST.swap(false, Ordering::Relaxed) {
2585            // Nothing dispatched yet.
2586        } else {
2587            FLUSHED.call_once(crate::gpu_wgpu::pipeline_cache_flush);
2588        }
2589    }
2590    GRAPH_RACE_TOK.store(0, Ordering::Relaxed);
2591    if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
2592        return;
2593    }
2594    let (gn, cn) = (
2595        GRAPH_N[1].load(Ordering::Relaxed),
2596        GRAPH_N[0].load(Ordering::Relaxed),
2597    );
2598    if gn >= GRAPH_RACE_SAMPLES && cn >= GRAPH_RACE_SAMPLES {
2599        let g_avg = GRAPH_NS[1].load(Ordering::Relaxed) / gn as u64;
2600        let c_avg = GRAPH_NS[0].load(Ordering::Relaxed) / cn as u64;
2601        let verdict = if g_avg < c_avg { 1 } else { 2 };
2602        GRAPH_RACE_STATE.store(verdict, Ordering::Relaxed);
2603        tracing::info!(
2604            "wgpu graph race: graph {:.2} ms/tok vs normal {:.2} ms/tok -> {}",
2605            g_avg as f64 / 1e6,
2606            c_avg as f64 / 1e6,
2607            if verdict == 1 { "graph" } else { "normal path" }
2608        );
2609        return;
2610    }
2611    let flip = GRAPH_RACE_FLIP.fetch_add(1, Ordering::Relaxed);
2612    GRAPH_RACE_ARM_GRAPH.store((flip % 2 == 1) as u8, Ordering::Relaxed);
2613}
2614
2615/// Should this decode token try the graph? `trusted` (discrete adapter,
2616/// explicit env, or a GDN hybrid whose state lives on the device) skips
2617/// the race entirely.
2618pub fn graph_race_use_graph(trusted: bool) -> bool {
2619    if trusted {
2620        return true;
2621    }
2622    match GRAPH_RACE_STATE.load(Ordering::Relaxed) {
2623        1 => true,
2624        2 => false,
2625        _ => GRAPH_RACE_ARM_GRAPH.load(Ordering::Relaxed) == 1,
2626    }
2627}
2628
2629/// First decode token of a racing graph generation: hopeless already?
2630/// (>4x the normal path's per-token average AND over a second.) Settles
2631/// the race immediately; the caller discards the graph result and
2632/// recomputes this token on the normal path.
2633pub fn graph_race_first_token_hopeless(dur: std::time::Duration) -> bool {
2634    if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
2635        return false;
2636    }
2637    let first = GRAPH_RACE_TOK.load(Ordering::Relaxed) == 0;
2638    let cn = GRAPH_N[0].load(Ordering::Relaxed);
2639    if !first || cn == 0 {
2640        return false;
2641    }
2642    let c_avg = GRAPH_NS[0].load(Ordering::Relaxed) / cn as u64;
2643    let ns = dur.as_nanos() as u64;
2644    if ns > 1_000_000_000 && ns > 4 * c_avg {
2645        GRAPH_RACE_STATE.store(2, Ordering::Relaxed);
2646        tracing::info!(
2647            "wgpu graph race: first graph token {:.0} ms vs normal {:.2} ms/tok — hopeless, normal path wins",
2648            ns as f64 / 1e6,
2649            c_avg as f64 / 1e6
2650        );
2651        return true;
2652    }
2653    false
2654}
2655
2656/// Record one decode-token wall time for the racing arm. The first
2657/// token of each generation is discarded (KV-mirror upload / cold
2658/// caches on the graph arm; cold mmap on the normal arm).
2659pub fn graph_race_record(used_graph: bool, dur: std::time::Duration) {
2660    if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
2661        return;
2662    }
2663    let tok = GRAPH_RACE_TOK.fetch_add(1, Ordering::Relaxed);
2664    if tok == 0 {
2665        return;
2666    }
2667    let i = used_graph as usize;
2668    GRAPH_NS[i].fetch_add(dur.as_nanos() as u64, Ordering::Relaxed);
2669    GRAPH_N[i].fetch_add(1, Ordering::Relaxed);
2670}
2671
2672/// Bounded-cost content fingerprint for the backends' pointer-keyed device
2673/// caches: FNV over the whole slice up to 4 KiB, over 64 spread 64-byte
2674/// windows (plus the length) above. An address-keyed hit must also prove
2675/// the bytes are still the ones it uploaded — the allocator reuses heap
2676/// and mmap addresses freely, so a reloaded model or a re-dequantized
2677/// layer lands where the old bytes were — and sampling keeps that proof at
2678/// ~a microsecond even for a 126 MB matrix. Real replacements (another
2679/// model's tensor, an Adam-updated master) differ densely, so a 4 KiB
2680/// spread cannot miss them.
2681pub(crate) fn fp_bytes(data: &[u8]) -> u64 {
2682    #[inline]
2683    fn fnv(mut h: u64, bytes: &[u8]) -> u64 {
2684        let (chunks, tail) = bytes.split_at(bytes.len() & !7);
2685        for c in chunks.chunks_exact(8) {
2686            h ^= u64::from_le_bytes(c.try_into().unwrap());
2687            h = h.wrapping_mul(0x100_0000_01b3);
2688        }
2689        for &b in tail {
2690            h ^= b as u64;
2691            h = h.wrapping_mul(0x100_0000_01b3);
2692        }
2693        h
2694    }
2695    let mut h = 0xcbf2_9ce4_8422_2325u64 ^ (data.len() as u64);
2696    if data.len() <= 4096 {
2697        return fnv(h, data);
2698    }
2699    let step = (data.len() - 64) / 63;
2700    for i in 0..64 {
2701        h = fnv(h, &data[i * step..i * step + 64]);
2702    }
2703    h
2704}
2705
2706/// `fp_bytes` over an f32 slice without a bytemuck dependency (the Metal
2707/// backend builds with no GPU feature flags).
2708pub(crate) fn fp_f32(data: &[f32]) -> u64 {
2709    let bytes = unsafe { std::slice::from_raw_parts(data.as_ptr() as *const u8, data.len() * 4) };
2710    fp_bytes(bytes)
2711}
2712
2713#[cfg(test)]
2714mod fp_tests {
2715    use super::fp_bytes;
2716
2717    /// The pointer-keyed caches survive on `fp_bytes` telling two different
2718    /// tensors apart at a reused address. Its sampling must therefore see a
2719    /// change ANYWHERE — head, tail, and the stretches between windows are
2720    /// the places a cheaper hash would go blind.
2721    #[test]
2722    fn fp_bytes_sees_a_change_anywhere_in_a_sampled_slice() {
2723        let n = 1 << 20; // 1 MiB — far above the 4 KiB full-hash threshold
2724        let base: Vec<u8> = (0..n).map(|i| (i * 31 + 7) as u8).collect();
2725        let h0 = fp_bytes(&base);
2726        assert_eq!(h0, fp_bytes(&base), "fingerprint must be deterministic");
2727        // A DENSE change (every requantized/redequantized tensor is one)
2728        // must flip the fingerprint no matter how the windows fall.
2729        let mut dense = base.clone();
2730        for b in dense.iter_mut() {
2731            *b = b.wrapping_add(1);
2732        }
2733        assert_ne!(
2734            h0,
2735            fp_bytes(&dense),
2736            "a fully different tensor slipped through"
2737        );
2738        // Length participates: the same prefix at a shorter length is a
2739        // different key AND a different fingerprint.
2740        assert_ne!(h0, fp_bytes(&base[..n - 64]));
2741        // Below the threshold the hash is exact: a single flipped byte in
2742        // a norm-sized vector must be seen.
2743        let mut small = vec![3u8; 4096];
2744        let hs = fp_bytes(&small);
2745        small[2048] ^= 1;
2746        assert_ne!(hs, fp_bytes(&small), "full hash missed a one-byte change");
2747        // And the sampled windows land within bounds on awkward sizes.
2748        for n in [4097usize, 5000, 64 * 64, 1 << 16] {
2749            let v = vec![9u8; n];
2750            let _ = fp_bytes(&v); // must not panic on window math
2751        }
2752    }
2753}
2754
2755/// Hand the card back after a bake: drop its resident weights, planes and
2756/// pools so the ordinary engine (the runtime gate, a serve that follows)
2757/// starts from a clean budget. No-op off the wgpu backend.
2758pub fn bake_release() {
2759    #[cfg(feature = "gpu")]
2760    crate::gpu_wgpu::bake_release();
2761}
2762
2763/// Strict-f32 for the bake's GEMMs (phase A mask training): the mask
2764/// selects neurons by a gradient signal, and f16 operand rounding on
2765/// that signal closes the wrong ones. No-op off the wgpu backend.
2766pub fn bake_precision_strict(on: bool) {
2767    #[cfg(feature = "gpu")]
2768    crate::gpu_wgpu::bake_precision_strict(on);
2769    #[cfg(not(feature = "gpu"))]
2770    let _ = on;
2771}
2772
2773/// CMF_GRAPH_HOSTPROF=1: how a graph token's wall splits between the
2774/// host encoding the command stream and the tail the GPU still owes
2775/// after encode. Fifteen GPU-side suspects measured null while the
2776/// bench counted 17.7k allocations a token — this is the instrument
2777/// that says whether the thief was on the host all along.
2778pub fn hostprof_encode_done(t0: std::time::Instant) {
2779    use std::sync::atomic::{AtomicU64, Ordering};
2780    static ENC: AtomicU64 = AtomicU64::new(0);
2781    static N: AtomicU64 = AtomicU64::new(0);
2782    if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
2783        return;
2784    }
2785    ENC.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
2786    let n = N.fetch_add(1, Ordering::Relaxed) + 1;
2787    if n % 100 == 0 {
2788        eprintln!(
2789            "hostprof: encode {:.2} ms/token over {n} tokens",
2790            ENC.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
2791        );
2792    }
2793}
2794
2795pub fn hostprof_total(t0: std::time::Instant) {
2796    use std::sync::atomic::{AtomicU64, Ordering};
2797    static TOT: AtomicU64 = AtomicU64::new(0);
2798    static N: AtomicU64 = AtomicU64::new(0);
2799    if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
2800        return;
2801    }
2802    TOT.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
2803    let n = N.fetch_add(1, Ordering::Relaxed) + 1;
2804    if n % 100 == 0 {
2805        eprintln!(
2806            "hostprof: total {:.2} ms/token over {n} tokens",
2807            TOT.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
2808        );
2809    }
2810}
2811
2812/// Per-stage host-encode accumulator for the Metal token loop
2813/// (CMF_GRAPH_HOSTPROF=1). Stage 0 = GDN-run encode; everything else
2814/// falls out by subtraction from hostprof's encode total.
2815pub fn stageprof(stage: u32, dt: std::time::Duration) {
2816    use std::sync::atomic::{AtomicU64, Ordering};
2817    static NS: [AtomicU64; 4] = [
2818        AtomicU64::new(0),
2819        AtomicU64::new(0),
2820        AtomicU64::new(0),
2821        AtomicU64::new(0),
2822    ];
2823    static N: AtomicU64 = AtomicU64::new(0);
2824    if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
2825        return;
2826    }
2827    NS[stage as usize % 4].fetch_add(dt.as_nanos() as u64, Ordering::Relaxed);
2828    if stage == 1 {
2829        let n = N.fetch_add(1, Ordering::Relaxed) + 1;
2830        if n % 200 == 0 {
2831            eprintln!(
2832                "stageprof: planning {:.2} ms/tok | gdn-item {:.2} ms/tok | attn-item {:.2} ms/tok ({n} tok)",
2833                NS[1].load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
2834                NS[2].load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
2835                NS[3].load(Ordering::Relaxed) as f64 / n as f64 / 1e6
2836            );
2837        }
2838    }
2839}
2840
2841/// Active weight bytes dispatched so far (Metal decode path); 0 where
2842/// the backend does not count. The honest floor's numerator.
2843pub fn weight_bytes_dispatched() -> u64 {
2844    let mut total = 0u64;
2845    #[cfg(target_os = "macos")]
2846    {
2847        total += crate::gpu_metal::WEIGHT_BYTES.load(std::sync::atomic::Ordering::Relaxed);
2848    }
2849    #[cfg(feature = "gpu")]
2850    {
2851        total += crate::gpu_wgpu::WEIGHT_BYTES.load(std::sync::atomic::Ordering::Relaxed);
2852    }
2853    total
2854}
2855
2856/// The per-stage split of `weight_bytes_dispatched`:
2857/// [misc, dense-ffn, moe, attn, gdn, head].
2858pub fn weight_bytes_by() -> [u64; 6] {
2859    #[cfg(target_os = "macos")]
2860    {
2861        let mut o = [0u64; 6];
2862        for (i, a) in crate::gpu_metal::WEIGHT_BYTES_BY.iter().enumerate() {
2863            o[i] = a.load(std::sync::atomic::Ordering::Relaxed);
2864        }
2865        return o;
2866    }
2867    #[allow(unreachable_code)]
2868    [0; 6]
2869}
2870
2871#[cfg(test)]
2872mod probe_warmup_tests {
2873    use super::*;
2874    use std::time::Duration;
2875
2876    fn ms(v: f64) -> Duration {
2877        Duration::from_nanos((v * 1e6) as u64)
2878    }
2879
2880    /// The bug this pins, measured on an A100: the first device call for
2881    /// a class compiles its pipeline, was timed at 117.01 ms against the
2882    /// host's 3.19, and sent `gemm-nt` to the CPU for the whole process —
2883    /// which ran a 27B bake on 2.6 cores with the card idle.
2884    #[test]
2885    fn one_cold_first_sample_does_not_lose_the_class() {
2886        let p = Probe::new();
2887        // First device sample is the pipeline build. Then the truth.
2888        probe_record_into(&p, "gemm-nt", None, true, ms(117.01));
2889        probe_record_into(&p, "gemm-nt", None, true, ms(1.1));
2890        probe_record_into(&p, "gemm-nt", None, true, ms(1.0));
2891        probe_record_into(&p, "gemm-nt", None, false, ms(3.19));
2892        probe_record_into(&p, "gemm-nt", None, false, ms(3.20));
2893        assert_eq!(
2894            p.state.load(Ordering::Relaxed),
2895            1,
2896            "the device is 3x faster once warm and must win"
2897        );
2898    }
2899
2900    /// The warm-up must not become a way to never decide, and must not
2901    /// underflow: a blind decrement at zero wraps a u32 to its maximum
2902    /// and mutes the arm for the life of the process.
2903    #[test]
2904    fn the_warmup_is_spent_once_and_never_underflows() {
2905        let p = Probe::new();
2906        for _ in 0..8 {
2907            probe_record_into(&p, "matmat", None, true, ms(10.0));
2908        }
2909        assert_eq!(p.gpu_burn.load(Ordering::Relaxed), 0, "spent, not wrapped");
2910        assert_eq!(
2911            p.gpu_n.load(Ordering::Relaxed),
2912            7,
2913            "one sample burned, the rest counted"
2914        );
2915    }
2916
2917    /// A device path that always refuses records no timing, so without
2918    /// counting the refusals the class can never reach a verdict. On an
2919    /// M4 with LFM2.5-2.6B `ffn` was still undecided after 9000 calls,
2920    /// alternating arms and paying a failed device attempt on half of
2921    /// them.
2922    #[test]
2923    fn a_class_whose_device_always_declines_settles_on_the_host() {
2924        // A class no other test in this file touches: `probe_note_decline`
2925        // works on the process-wide probes by design, and the tests in
2926        // this binary share them.
2927        let c = OpClass::MatmatWide;
2928        let p = &PROBES[c as usize];
2929        p.state.store(0, Ordering::Relaxed);
2930        p.declines.store(0, Ordering::Relaxed);
2931        for _ in 0..(PROBE_DECLINE_LIMIT - 1) {
2932            probe_note_decline(c);
2933        }
2934        assert_eq!(
2935            p.state.load(Ordering::Relaxed),
2936            0,
2937            "one short of the limit is still a question, not an answer"
2938        );
2939        probe_note_decline(c);
2940        assert_eq!(p.state.load(Ordering::Relaxed), 2, "settled on the host");
2941        assert!(matches!(probe_arm(c), ProbeArm::Cpu));
2942        p.state.store(0, Ordering::Relaxed);
2943        p.declines.store(0, Ordering::Relaxed);
2944    }
2945
2946    /// A genuinely slower device still loses — the warm-up removes an
2947    /// artefact, it does not put a thumb on the scale.
2948    #[test]
2949    fn a_slow_device_still_loses_after_the_warmup() {
2950        let p = Probe::new();
2951        for _ in 0..4 {
2952            probe_record_into(&p, "matvec", None, true, ms(40.0));
2953        }
2954        for _ in 0..4 {
2955            probe_record_into(&p, "matvec", None, false, ms(2.0));
2956        }
2957        assert_eq!(p.state.load(Ordering::Relaxed), 2, "host wins on merit");
2958    }
2959}