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
18/// Packed, f32-only Embryo model consumed by the resident Vulkan graph.
19///
20/// This is deliberately a separate representation from `CmfModel`: the
21/// Embryo graph owns one contiguous device copy and never dereferences the
22/// mmap or asks the ordinary per-op residency path for a tensor.  `meta`
23/// contains byte-free element offsets into `weights`; `UINT_MAX` marks an
24/// absent optional matrix.  The builder lives in `pipeline.rs`, where the
25/// architecture-specific weight types are visible.
26pub struct EmbryoGraphModel {
27    pub id: u64,
28    pub hidden: usize,
29    pub intermediate: usize,
30    pub vocab: usize,
31    pub layers: usize,
32    pub phase_heads: usize,
33    pub nphase: usize,
34    pub phase_dv: usize,
35    pub anchor_q_heads: usize,
36    pub anchor_kv_heads: usize,
37    pub anchor_head_dim: usize,
38    pub rotary_dim: usize,
39    pub max_seq: usize,
40    pub cluster_count: usize,
41    pub cluster_size: usize,
42    pub phase_state_len: usize,
43    pub state_stride: usize,
44    pub kv_stride: usize,
45    pub norm_gemma: bool,
46    pub phase_mass: f32,
47    pub weights: Vec<f32>,
48    pub meta: Vec<u32>,
49    pub lm_head: Vec<f32>,
50    pub clusters: Vec<f32>,
51    pub final_norm: Vec<f32>,
52    pub inv_freq: Vec<f32>,
53    /// Natively bounded anchors (`swa_sink_v1`): every anchor layer owns a
54    /// ring of `anchor_window` raw keys/values instead of a `max_seq` KV
55    /// plane, so `kv_stride` is the ring size, `kv_layers` counts the
56    /// anchors, and no position cap applies (the operator has none).
57    pub bounded: bool,
58    /// Layers that own a KV/ring slot of `kv_stride` f32 (all layers for
59    /// the legacy full anchor — offsets are `layer·kv_stride` — or the
60    /// bounded anchors only).
61    pub kv_layers: usize,
62    /// Recurrent (phase or GDN) layers — the state buffer holds one
63    /// `state_stride` per such layer; anchors own none.
64    pub state_layers: usize,
65    pub anchor_window: usize,
66    pub anchor_sink: usize,
67    /// vmf_phase mixer layers (kind 0/1) in the stack.
68    pub phase_layers: usize,
69    /// GatedDeltaNet mixer layers (kind 4) in the stack; their geometry
70    /// is `gdn_heads` (nv), `gdn_k_heads` (nk), `gdn_dk`, `gdn_dv`,
71    /// `gdn_kk` (conv taps) — header words 24..29 of `meta`.
72    pub gdn_layers: usize,
73    pub gdn_heads: usize,
74    pub gdn_k_heads: usize,
75    pub gdn_dk: usize,
76    pub gdn_dv: usize,
77    pub gdn_kk: usize,
78}
79
80impl EmbryoGraphModel {
81    /// Fused q/k/v projection width of the GDN mixer (`2·nk·dk + nv·dv`).
82    pub fn gdn_c_dim(&self) -> usize {
83        2 * self.gdn_k_heads * self.gdn_dk + self.gdn_heads * self.gdn_dv
84    }
85}
86
87/// Words in the resident graph's `meta` header before the per-layer
88/// records (64 words each). Shared by the host packer and the encoder.
89pub const EMBRYO_META_HEADER: usize = 32;
90
91/// Positions one chunked-prefill submit of the resident Embryo graph
92/// covers (the shader's `CMAX`; the chunk window of the scratch buffer is
93/// sized by it).
94pub const EMBRYO_CHUNK_MAX: usize = 64;
95
96/// `n` contiguous prompt positions (`rows` = n × hidden embeddings from
97/// `position`) through the resident Embryo graph in one submit; logits of
98/// the last position. False = refused before any device work (the caller
99/// keeps the per-position path), or the sequence went cold.
100pub fn forward_embryo_graph_chunk(
101    model: &Arc<EmbryoGraphModel>,
102    kv_id: u64,
103    rows: &[f32],
104    position: usize,
105    n: usize,
106    logits: &mut Vec<f32>,
107) -> bool {
108    #[cfg(feature = "gpu")]
109    if backend() == Backend::Wgpu {
110        return crate::gpu_wgpu::forward_embryo_graph_chunk(model, kv_id, rows, position, n, logits);
111    }
112    let _ = (model, kv_id, rows, position, n, logits);
113    false
114}
115
116/// Bytes the resident Embryo graph holds on the device for a sequence:
117/// `(recurrent state, anchor KV/ring)`. None when the sequence has no
118/// device image (host path, or not started yet). This is the measured
119/// long-context claim of the resident path — a bounded genome reports
120/// the same two numbers at every context depth.
121pub fn embryo_device_state_bytes(kv_id: u64) -> Option<(u64, u64)> {
122    #[cfg(feature = "gpu")]
123    if backend() == Backend::Wgpu {
124        return crate::gpu_wgpu::embryo_device_state_bytes(kv_id);
125    }
126    let _ = kv_id;
127    None
128}
129
130/// Does this directory entry count as WEIGHT in the device placement
131/// heuristics (layer prefix, live-weight budget, `serve` replicas vs
132/// split)? On a genome file only the trunk does: skill tensors replace
133/// trunk tensors of the same shape inside their own lane, so counting
134/// them would give the backbone of F1 another GPU prefix than F0's
135/// (NF-7). Other files keep the historical whole-directory count.
136pub fn counts_as_placement_weight(model: &CmfModel, name: &str) -> bool {
137    model.header.genome.is_none() || cortiq_core::knowledge::is_trunk_tensor(name)
138}
139
140/// Weight bytes the placement heuristics budget for `model`: the trunk's
141/// payload on a genome file ([`counts_as_placement_weight`]), the whole
142/// mapped file otherwise (the historical measure).
143pub fn placement_weight_bytes(model: &CmfModel) -> u64 {
144    if model.header.genome.is_none() {
145        return model.primary_bytes().len() as u64;
146    }
147    model
148        .tensors
149        .iter()
150        .filter(|t| counts_as_placement_weight(model, &t.name))
151        .map(|t| t.nbytes)
152        .sum()
153}
154
155/// Next position of the resident Embryo sequence `kv_id` (`Some(n)` = the
156/// device holds positions `0..n` of it), `None` when the device holds no
157/// initialized image of that sequence — the host owns it (or none began).
158pub fn embryo_device_next_position(kv_id: u64) -> Option<usize> {
159    #[cfg(feature = "gpu")]
160    if backend() == Backend::Wgpu {
161        return crate::gpu_wgpu::embryo_device_next_position(kv_id);
162    }
163    let _ = kv_id;
164    None
165}
166
167thread_local! {
168    /// Index of the current forward layer (−1 = outside a numbered layer:
169    /// lm_head/embed — always allowed). The pipeline sets it before
170    /// each layer so that the GPU/CPU layer-split works.
171    static CUR_LAYER: Cell<i64> = const { Cell::new(-1) };
172    /// Inside `cpu_scope` every GPU gate reports disabled: the timed CPU
173    /// arm of a probe (and a class that lost its probe) must run PURE
174    /// CPU, or inner per-op hooks would re-enter the GPU and poison the
175    /// comparison.
176    static CPU_ONLY: Cell<bool> = const { Cell::new(false) };
177    /// "This op paid a one-off cost" (weight upload / first pipeline
178    /// build): backends set it, `probe_record` discards the sample so
179    /// only steady-state timings compete.
180    static PROBE_COLD: Cell<bool> = const { Cell::new(false) };
181}
182
183/// RAII form of `cpu_scope`, used when a device-prefix graph hands a whole
184/// remainder of the forward pass to the host. Without a guard around that
185/// tail, its ordinary per-op hooks re-entered the GPU and streamed the rest of
186/// an over-size model through the residency arena, defeating the prefix's VRAM
187/// bound at the driver-allocation level.
188pub struct CpuScopeGuard(bool);
189
190impl Drop for CpuScopeGuard {
191    fn drop(&mut self) {
192        CPU_ONLY.with(|c| c.set(self.0));
193    }
194}
195
196pub fn enter_cpu_scope() -> CpuScopeGuard {
197    let previous = CPU_ONLY.with(|c| c.replace(true));
198    CpuScopeGuard(previous)
199}
200
201/// Run `f` with the GPU gates off on this thread (pure-CPU arm).
202pub fn cpu_scope<R>(f: impl FnOnce() -> R) -> R {
203    let _restore = enter_cpu_scope();
204    f()
205}
206
207/// Capture CPU-only placement before dispatching whole operators to workers.
208/// `cpu_scope` is thread-local, while a MoE panel worker calls QTensor again;
209/// without inheritance, which expert happened to land on the caller changed
210/// its precision/backend from run to run.
211pub(crate) fn inherit_cpu_scope() -> impl Fn() -> Option<CpuScopeGuard> + Copy {
212    let on = CPU_ONLY.get();
213    move || on.then(enter_cpu_scope)
214}
215
216/// Backends: name the device once at init. The probe cache is keyed by
217/// it, because a verdict is a property of THIS silicon and nothing else.
218/// First writer wins: a process runs one backend, and on the rare host
219/// where two initialize, the one that came up first is the one in use.
220pub fn probe_set_device(label: &str) {
221    let _ = DEVICE_LABEL.set(label.to_string());
222}
223
224fn device_label() -> &'static str {
225    DEVICE_LABEL.get().map(String::as_str).unwrap_or("unknown")
226}
227
228static DEVICE_LABEL: std::sync::OnceLock<String> = std::sync::OnceLock::new();
229
230/// Somewhere this process may write small caches.
231///
232/// `std::env::temp_dir()` is NOT that place on Android: with no `TMPDIR`
233/// it answers `/tmp`, which does not exist in an app sandbox, and every
234/// write fails silently — measured, after the pipeline cache appeared to
235/// work in a shell (where `TMPDIR=/data/local/tmp`) and did nothing at
236/// all in the app. The loader points this at the model's own directory,
237/// which is somewhere the caller already writes.
238static CACHE_DIR: std::sync::OnceLock<std::path::PathBuf> = std::sync::OnceLock::new();
239
240/// Loader: name a directory this process can write to. First call wins.
241pub fn set_cache_dir(dir: std::path::PathBuf) {
242    let _ = CACHE_DIR.set(dir);
243}
244
245/// Same directory, for the backends.
246pub fn cache_dir_pub() -> std::path::PathBuf {
247    cache_dir()
248}
249
250fn cache_dir() -> std::path::PathBuf {
251    if let Some(d) = CACHE_DIR.get() {
252        return d.clone();
253    }
254    match std::env::var_os("TMPDIR") {
255        Some(t) => std::path::PathBuf::from(t),
256        None => std::env::temp_dir(),
257    }
258}
259
260/// Where decided verdicts are remembered between runs. `CMF_PROBE_CACHE`
261/// overrides the path; `0` disables the cache entirely.
262fn probe_cache_path() -> Option<std::path::PathBuf> {
263    match std::env::var("CMF_PROBE_CACHE") {
264        Ok(v) if v == "0" => None,
265        Ok(v) => Some(std::path::PathBuf::from(v)),
266        Err(_) => Some(cache_dir().join("cortiq-gpu-probe.tsv")),
267    }
268}
269
270/// One line per decided class: `version \t device \t class \t winner`.
271/// A different engine build or a different device simply does not match,
272/// so a stale file is inert rather than wrong.
273fn probe_cache_key_named(class: &str) -> String {
274    format!(
275        "{}\t{}\t{}",
276        env!("CARGO_PKG_VERSION"),
277        device_label(),
278        class
279    )
280}
281
282const CLASS_NAMES: [&str; 7] = [
283    "ffn",
284    "matvec",
285    "matmat",
286    "qkv-batch",
287    "matmat-wide",
288    "lm-head",
289    "gemm-nt",
290];
291
292/// Adopt every verdict this device already reached in an earlier run.
293///
294/// Probing is not cheap and it is not free of consequences: on a
295/// Snapdragon 778G the three deciding classes took **three minutes of
296/// wall clock** before the first token, every process, and in the phone
297/// app that was the whole first answer — 209.6 s for 25 tokens against
298/// 10.5 s on the CPU path. The verdict itself was the same every time.
299/// Paying to rediscover it is the defect; the answer is to write it down.
300fn probe_cache_load() {
301    static ONCE: std::sync::Once = std::sync::Once::new();
302    ONCE.call_once(|| {
303        let Some(path) = probe_cache_path() else {
304            return;
305        };
306        // Unit tests share this process and its default cache path; a
307        // verdict left by an earlier run would decide a class before the
308        // arbitration tests get to watch it alternate. Tests that mean to
309        // exercise the cache point `CMF_PROBE_CACHE` at their own file.
310        if cfg!(test) && std::env::var("CMF_PROBE_CACHE").is_err() {
311            return;
312        }
313        let Ok(text) = std::fs::read_to_string(&path) else {
314            return;
315        };
316        probe_cache_adopt(&text);
317    });
318}
319
320/// Apply verdicts from a cache file's text. Split out from the file
321/// reading so the adoption rule — including which lines must be IGNORED
322/// — is testable without a filesystem.
323fn probe_cache_adopt(text: &str) {
324    for line in text.lines() {
325        let Some((key, verdict)) = line.rsplit_once('\t') else {
326            continue;
327        };
328        let winner = match verdict.trim() {
329            "gpu" => 1u8,
330            "cpu" => 2u8,
331            _ => continue,
332        };
333        for (i, name) in CLASS_NAMES.iter().enumerate() {
334            if probe_cache_key_named(name) == key {
335                let _ = PROBES[i].state.compare_exchange(
336                    0,
337                    winner,
338                    Ordering::Relaxed,
339                    Ordering::Relaxed,
340                );
341                tracing::debug!("gpu probe [{name}]: remembered → {verdict}");
342            }
343        }
344    }
345}
346
347/// Remember a verdict for the next run. Best-effort: a read-only cache
348/// directory costs a re-probe, never a failure.
349fn probe_cache_store(c: OpClass, winner: u8) {
350    let Some(path) = probe_cache_path() else {
351        return;
352    };
353    let line = format!(
354        "{}\t{}\n",
355        probe_cache_key_named(CLASS_NAMES[c as usize]),
356        if winner == 1 { "gpu" } else { "cpu" }
357    );
358    use std::io::Write;
359    if let Ok(mut f) = std::fs::OpenOptions::new()
360        .create(true)
361        .append(true)
362        .open(&path)
363    {
364        let _ = f.write_all(line.as_bytes());
365    }
366}
367
368/// Backends: note a one-off cost (weight upload, buffer-cache fill) so
369/// the probe discards this sample.
370/// Every buffer creation anywhere bumps this; the graph's bind-group
371/// cache treats any cold event as total invalidation — a stale bind
372/// group is silent corruption, a cleared cache is one re-encoded token.
373pub fn cold_epoch() -> u64 {
374    COLD_EPOCH.load(std::sync::atomic::Ordering::Relaxed)
375}
376static COLD_EPOCH: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
377
378pub(crate) fn probe_note_cold() {
379    COLD_EPOCH.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
380    PROBE_COLD.with(|c| c.set(true));
381}
382
383/// Peek the cold flag without consuming it (`probe_record` consumes).
384/// Contention heuristics use this: a slow COLD op is a one-off build
385/// cost, not evidence the device is busy.
386pub(crate) fn probe_was_cold() -> bool {
387    PROBE_COLD.with(|c| c.get())
388}
389
390/// Pipeline: mark the current layer (or −1 outside layers) for layer-split.
391pub fn set_layer(l: i64) {
392    CUR_LAYER.with(|c| c.set(l));
393}
394
395/// The layer `set_layer` last marked on this thread (−1 outside layers).
396pub fn cur_layer() -> i64 {
397    CUR_LAYER.with(|c| c.get())
398}
399
400/// Capacity-derived layer prefix for per-op walks. The explicit
401/// `CMF_GPU_LAYERS` override is handled by the backend and takes precedence.
402pub fn automatic_layer_prefix(
403    model: &Arc<CmfModel>,
404    num_layers: usize,
405    physical_layers: usize,
406) -> Option<usize> {
407    match backend() {
408        #[cfg(feature = "gpu")]
409        Backend::Wgpu => {
410            crate::gpu_wgpu::automatic_layer_prefix(model, num_layers, physical_layers)
411        }
412        _ => None,
413    }
414}
415
416/// Parse `CMF_GPU_LAYERS` («0-19», «0,2,4», «0-9,30-39») once.
417/// None = no restriction (all layers on GPU). Garbage → also no restriction.
418fn layer_ranges() -> &'static Option<Vec<(i64, i64)>> {
419    static R: OnceLock<Option<Vec<(i64, i64)>>> = OnceLock::new();
420    R.get_or_init(|| {
421        let s = std::env::var("CMF_GPU_LAYERS").ok()?;
422        let mut v = Vec::new();
423        for part in s.split(',') {
424            let part = part.trim();
425            match part.split_once('-') {
426                Some((a, b)) => v.push((a.trim().parse().ok()?, b.trim().parse().ok()?)),
427                None => {
428                    let x: i64 = part.parse().ok()?;
429                    v.push((x, x));
430                }
431            }
432        }
433        Some(v)
434    })
435}
436
437fn layer_allowed() -> bool {
438    match layer_ranges() {
439        None => true,
440        Some(ranges) => {
441            let cur = CUR_LAYER.with(|c| c.get());
442            cur < 0 || ranges.iter().any(|(a, b)| cur >= *a && cur <= *b)
443        }
444    }
445}
446
447/// GPU allowed FOR THE CURRENT LAYER: backend is initialized AND the layer
448/// falls within `CMF_GPU_LAYERS` (GPU/CPU layer-split) AND we are not
449/// inside a `cpu_scope`. Op gates call this.
450pub fn enabled_here() -> bool {
451    !CPU_ONLY.with(|c| c.get()) && enabled() && layer_allowed()
452}
453
454/// Descriptor-aware q2tp Vulkan kernels are kept behind an explicit opt-in
455/// until the Prism full-graph/resident-weight path has a coherent generation
456/// gate.  `CMF_GPU=1` alone must not silently turn synchronous per-op
457/// readbacks into the default model path; callers and validation tests can
458/// request the measured kernels with `CMF_Q2TP_GPU=1`.
459pub fn q2tp_gpu_opt_in() -> bool {
460    std::env::var("CMF_Q2TP_GPU").as_deref() == Ok("1")
461}
462
463// ── Runtime GPU-vs-CPU probe ────────────────────────────────────────────
464// CMF_GPU=1 does not TRUST that the device wins — it MEASURES. For each
465// op class the first calls alternate arms: GPU timed vs pure-CPU timed
466// (under cpu_scope). Cold GPU calls (weight upload / cache fill) are
467// discarded; after PROBE_SAMPLES clean samples per arm the faster arm is
468// chosen for the rest of the process. Rationale: submit+poll latency
469// differs by an order of magnitude across driver stacks (Metal/PCIe
470// ~3-4 ms, Vulkan/4090 ~0.3 ms) — a static threshold cannot know whether
471// per-op offload pays off HERE. CMF_GPU_PROBE=0 → always trust the GPU.
472
473/// GPU-eligible op classes, each with an independent probe.
474#[derive(Clone, Copy)]
475pub enum OpClass {
476    /// Whole FFN chain in one submission (dense / MoE block).
477    Ffn = 0,
478    /// Large hybrid CPU∥GPU matvec (lm_head class).
479    Matvec = 1,
480    /// Prefill GEMM (matmat).
481    Matmat = 2,
482    /// Batched matvecs of one input (QKV).
483    Batch = 3,
484    /// Prefill GEMM at image-diffusion widths (b ≥ 128). Probed apart
485    /// from `Matmat`: one imagegen process runs BOTH populations
486    /// (prompt encode b≈40 where the GPU wins big, DiT b≥256 where
487    /// the CPU AMX arm is competitive) — a single shared verdict locks
488    /// the wrong arm for whichever population samples second.
489    MatmatWide = 4,
490    /// The lm_head itself, apart from the merely-large matvecs. Same
491    /// reasoning as `MatmatWide`, and DeepSeek-V4 is where it bit: its
492    /// attention projections are 37M weights and its head is 529M, so
493    /// the projections' verdict — CPU, honestly measured at 0.19 ms —
494    /// decided for a matvec fourteen times their size that took 11 ms
495    /// a token on the host.
496    MatvecHead = 5,
497    /// The blocked f32 GEMM (`fcd_ops::gemm_nt`): attention's QKᵀ and
498    /// AV, and the VAE decoders' projections. It used to take every job
499    /// over 4 M MACs on sight, with no CPU arm to lose to — which on
500    /// the MiniMax-H3 video decoder was three times SLOWER than the
501    /// host it displaced. Its population is per-head slices, nothing
502    /// like the weight GEMMs above, so it probes on its own.
503    GemmNt = 6,
504}
505
506/// Which probe a large matvec belongs to. The head is an order of
507/// magnitude bigger than anything else that reaches this gate, and the
508/// two populations do not have the same answer.
509pub fn matvec_class(rows: usize, cols: usize) -> OpClass {
510    if rows * cols >= 67_108_864 {
511        OpClass::MatvecHead
512    } else {
513        OpClass::Matvec
514    }
515}
516
517/// Probe verdict for one call.
518pub enum ProbeArm {
519    /// Run the GPU path (during probing: timed, recorded).
520    Gpu,
521    /// Probing: run the CPU path under `cpu_scope`, timed, recorded.
522    CpuTimed,
523    /// Decided: CPU won — run the CPU path (under `cpu_scope`).
524    Cpu,
525}
526
527/// Clean samples per arm before a class decides.
528const PROBE_SAMPLES: u32 = 6;
529
530/// Declines before a class gives the work to the host for good. High
531/// enough that a transient refusal — an unsealed state during prefill, a
532/// shape the kernel skips this once — cannot settle the question.
533const PROBE_DECLINE_LIMIT: u32 = 16;
534
535/// Device samples discarded before any count — see `Probe::gpu_burn`.
536const PROBE_WARMUP: u32 = 1;
537
538struct Probe {
539    /// 0 = probing, 1 = GPU won, 2 = CPU won.
540    state: AtomicU8,
541    flip: AtomicU32,
542    gpu_ns: AtomicU64,
543    gpu_n: AtomicU32,
544    /// Times the device arm was chosen and the device DECLINED.
545    ///
546    /// A decline carries no timing, so nothing is recorded — and a class
547    /// whose device path always refuses therefore never reaches a
548    /// verdict, alternates arms forever, and pays a failed device
549    /// attempt on half of every token's calls. Measured on an M4 with
550    /// LFM2.5-2.6B: `ffn` was still undecided after 9000 calls, and a
551    /// token cost 83.55 ms against 41.85 with the device off — twice the
552    /// price for work the host did anyway.
553    declines: AtomicU32,
554    /// GPU samples still to discard as warm-up.
555    ///
556    /// The cold flag catches buffer and weight uploads, but a compute
557    /// pipeline is compiled on first use and not every creation site
558    /// raises it — the wgpu path has 21 pipeline creations against 12
559    /// cold notes. One uncaught shader compile is enough to lose a
560    /// class for the whole process: `gemm-nt` on an A100 was recorded at
561    /// 117.01 ms against the host's 3.19 and sent to the CPU, which
562    /// parked a 27B bake on 2.6 cores with the card idle. The decision
563    /// already uses each arm's BEST sample, so discarding the first
564    /// GPU sample costs one extra round trip and removes the whole
565    /// class of first-call artefacts.
566    gpu_burn: AtomicU32,
567    cpu_ns: AtomicU64,
568    cpu_n: AtomicU32,
569    /// Best (minimum) sample per arm. The DECISION compares these:
570    /// means are poisoned by one-off cold costs the cold-flag cannot
571    /// see — e.g. the CPU arm's first mmap-cold expert matvec page
572    /// faults its weights in and reads 3× its steady state, which
573    /// locked the GPU arm on a 35B MoE at a 4× real-world loss. The
574    /// minimum is each arm's honest steady-state pace.
575    gpu_min: AtomicU64,
576    cpu_min: AtomicU64,
577}
578
579impl Probe {
580    const fn new() -> Self {
581        Self {
582            state: AtomicU8::new(0),
583            flip: AtomicU32::new(0),
584            gpu_ns: AtomicU64::new(0),
585            gpu_n: AtomicU32::new(0),
586            declines: AtomicU32::new(0),
587            gpu_burn: AtomicU32::new(PROBE_WARMUP),
588            cpu_ns: AtomicU64::new(0),
589            cpu_n: AtomicU32::new(0),
590            gpu_min: AtomicU64::new(u64::MAX),
591            cpu_min: AtomicU64::new(u64::MAX),
592        }
593    }
594}
595
596static PROBES: [Probe; 7] = [
597    Probe::new(),
598    Probe::new(),
599    Probe::new(),
600    Probe::new(),
601    Probe::new(),
602    Probe::new(),
603    Probe::new(),
604];
605
606/// A caller that knows its loop is long, uniform and warm can say so: the
607/// probe times ops in isolation and alternates arms to do it, which reads a
608/// sustained diffusion step as slower on the device than it is. Measured on
609/// an M4 at 672 video tokens: the probe picked the CPU at 1.25 ms against
610/// 0.88 ms per op, and the loop it picked for ran 23.9 s a step against the
611/// device's 19.7 s.
612static TRUST_GPU: AtomicBool = AtomicBool::new(false);
613
614/// Take the probe out of the loop until the guard drops.
615pub fn trust_gpu() -> GpuTrust {
616    let was = TRUST_GPU.swap(true, Ordering::Relaxed);
617    GpuTrust(was)
618}
619
620pub struct GpuTrust(bool);
621
622impl Drop for GpuTrust {
623    fn drop(&mut self) {
624        TRUST_GPU.store(self.0, Ordering::Relaxed);
625    }
626}
627
628fn probe_on_for(c: OpClass) -> bool {
629    // The trust is only for the *wide* class. A sustained diffusion step is
630    // where the probe reads a warm device as cold; the narrow batches inside
631    // the same loop — an audio stream of fifty-one tokens against the same
632    // weights — are small enough that submit latency can genuinely beat the
633    // arithmetic, and there the probe is right and should keep deciding.
634    if TRUST_GPU.load(Ordering::Relaxed) && matches!(c, OpClass::MatmatWide | OpClass::Ffn) {
635        return false;
636    }
637    probe_on()
638}
639
640/// Is the per-op GPU/CPU probe enabled (`CMF_GPU_PROBE`, default on)?
641/// The native Metal decode route never consults it — `q1_force` routes
642/// the token graph to the device outright — so it is reported, not used.
643pub fn probe_enabled() -> bool {
644    probe_on()
645}
646
647fn probe_on() -> bool {
648    static ON: OnceLock<bool> = OnceLock::new();
649    *ON.get_or_init(|| {
650        std::env::var("CMF_GPU_PROBE")
651            .map(|v| v != "0" && v != "off")
652            .unwrap_or(true)
653    })
654}
655
656/// q1 ops on the native Metal backend skip the probe entirely: the CPU
657/// q1 kernel is load-port-bound, the GPU one wins warm — and probe
658/// alternation itself cools the device between samples (measured: block
659/// times 5.8 ms warm vs 8.8 ms mixed). Other backends keep probing.
660pub fn q1_force() -> bool {
661    #[cfg(target_os = "macos")]
662    {
663        backend() == Backend::Metal
664    }
665    #[cfg(not(target_os = "macos"))]
666    {
667        false
668    }
669}
670
671/// Should a FUSED whole-block path trust the device instead of asking
672/// the per-op probe? True on native Metal and on discrete wgpu adapters.
673///
674/// The probe answers "is one wide matmat faster on the GPU", and for the
675/// DiT on Metal that is a coin flip — measured 2.62 ms GPU vs 2.56 ms
676/// CPU, a 2% spread that lands on either arm run to run. But the fused
677/// block's advantage is not per-op speed, it is that the hidden state,
678/// the packs and the attention panels never leave the device: end to end
679/// the whole-block path renders a 512² Lumina step in ~5.4 s against
680/// ~8.4 s when the probe happens to pick the CPU. Gating a fusion win on
681/// a per-op tie made every second render half-speed at random.
682///
683/// On a discrete card the verdict is never in doubt — an RTX 3090 against
684/// a 256-core EPYC measured 11.5 ms vs 31 ms per wide op, four runs out
685/// of four — so the probe's sampling phase is pure cost: it alone was 10%
686/// of a 512² render (74.3 s against 66.9 s with the probe off). Integrated
687/// and mobile adapters keep probing; there the submit latency is real and
688/// can genuinely lose.
689pub fn fused_block_trusted() -> bool {
690    #[cfg(target_os = "macos")]
691    if backend() == Backend::Metal {
692        return true;
693    }
694    wgpu_graph_default()
695}
696
697/// Which arm should this GPU-eligible call take? Consult AFTER the
698/// eligibility gates (`enabled_here` / `min_rows`) so only real
699/// candidates alternate.
700/// While a class is still probing, a call whose weights are NOT yet on
701/// the card should take the GPU arm anyway: the upload is work the next
702/// step needs regardless, and the sample it produces is discarded as
703/// cold — so handing that call to the CPU arm buys nothing and costs a
704/// host GEMM. Measured on a diffusion stack, where every layer is
705/// touched once per step and therefore EVERY first-step GPU sample is
706/// cold: one projection drew the CPU arm for the whole first step, 9.8 s
707/// against the 2.8 s it costs once the weights are warm.
708pub fn weight_is_resident(model: &Arc<CmfModel>, idx: usize) -> bool {
709    #[cfg(all(feature = "gpu", not(target_os = "macos")))]
710    {
711        return crate::gpu_wgpu::weight_is_resident(model, idx);
712    }
713    #[cfg(not(all(feature = "gpu", not(target_os = "macos"))))]
714    {
715        let _ = (model, idx);
716        true
717    }
718}
719
720pub fn probe_arm_cold_prefers_gpu(c: OpClass, weights_resident: bool) -> ProbeArm {
721    if !weights_resident && probe_deciding(c) {
722        return ProbeArm::Gpu;
723    }
724    probe_arm(c)
725}
726
727pub fn probe_arm(c: OpClass) -> ProbeArm {
728    // Every arbitrated call starts with a clean cold flag: both the
729    // sample discard in `probe_record` and the contention kill-switch
730    // read it AFTER the op, so a stale note from a previous call on
731    // this thread must not leak in.
732    PROBE_COLD.with(|f| f.set(false));
733    if !probe_on_for(c) {
734        return ProbeArm::Gpu;
735    }
736    probe_cache_load();
737    let p = &PROBES[c as usize];
738    match p.state.load(Ordering::Relaxed) {
739        1 => ProbeArm::Gpu,
740        2 => ProbeArm::Cpu,
741        _ => {
742            if p.flip.fetch_add(1, Ordering::Relaxed) % 2 == 0 {
743                ProbeArm::Gpu
744            } else {
745                ProbeArm::CpuTimed
746            }
747        }
748    }
749}
750
751/// The device arm was chosen and the device refused the work, so there
752/// is no time to record. Callers that fall through to the host MUST say
753/// so here, or the class can never decide.
754pub fn probe_note_decline(c: OpClass) {
755    let p = &PROBES[c as usize];
756    if p.state.load(Ordering::Relaxed) != 0 {
757        return;
758    }
759    let n = p.declines.fetch_add(1, Ordering::Relaxed) + 1;
760    if n >= PROBE_DECLINE_LIMIT
761        && p.state
762            .compare_exchange(0, 2, Ordering::Relaxed, Ordering::Relaxed)
763            .is_ok()
764    {
765        tracing::info!(
766            "gpu probe [{}]: device declined {n} times → cpu",
767            CLASS_NAMES[c as usize]
768        );
769    }
770}
771
772/// Record a timed arm sample; on the `PROBE_SAMPLES`-th clean sample of
773/// BOTH arms the class decides for the rest of the process.
774pub fn probe_record(c: OpClass, gpu: bool, dur: std::time::Duration) {
775    probe_record_into(
776        &PROBES[c as usize],
777        CLASS_NAMES[c as usize],
778        Some(c),
779        gpu,
780        dur,
781    )
782}
783
784/// The body of `probe_record` over ONE probe, so the decision can be
785/// driven in a test without touching the process-wide array.
786fn probe_record_into(
787    p: &Probe,
788    class_name: &str,
789    cache: Option<OpClass>,
790    gpu: bool,
791    dur: std::time::Duration,
792) {
793    if p.state.load(Ordering::Relaxed) != 0 {
794        return;
795    }
796    if gpu && PROBE_COLD.with(|f| f.replace(false)) {
797        return; // one-off cost in this call — not a steady-state sample
798    }
799    if gpu {
800        // Load-then-store rather than fetch_sub: a blind decrement at
801        // zero wraps a u32 to its maximum and mutes the arm forever.
802        // A benign race here burns one extra sample, which is free.
803        let left = p.gpu_burn.load(Ordering::Relaxed);
804        if left > 0 {
805            p.gpu_burn.store(left - 1, Ordering::Relaxed);
806            return; // warm-up: the first device sample builds its pipeline
807        }
808    }
809    let ns = dur.as_nanos().min(u64::MAX as u128) as u64;
810    if gpu {
811        p.gpu_ns.fetch_add(ns, Ordering::Relaxed);
812        p.gpu_n.fetch_add(1, Ordering::Relaxed);
813        p.gpu_min.fetch_min(ns, Ordering::Relaxed);
814    } else {
815        p.cpu_ns.fetch_add(ns, Ordering::Relaxed);
816        p.cpu_n.fetch_add(1, Ordering::Relaxed);
817        p.cpu_min.fetch_min(ns, Ordering::Relaxed);
818    }
819    let (gn, cn) = (
820        p.gpu_n.load(Ordering::Relaxed),
821        p.cpu_n.load(Ordering::Relaxed),
822    );
823    if gn >= 2 && cn >= 2 {
824        // Decide on each arm's BEST sample — the steady-state pace.
825        // Means carry one-off cold costs (mmap page-in on the CPU arm)
826        // that the cold-flag machinery cannot see.
827        let g = p.gpu_min.load(Ordering::Relaxed) as f64;
828        let cp = p.cpu_min.load(Ordering::Relaxed) as f64;
829        // Early verdict on a ≥2× gap — no reason to keep feeding the
830        // losing arm; close races take the full sample count. It was 3×,
831        // and the cost of that half-octave was measured: a DiT whose
832        // wide GEMMs run 11.4 ms on the device against 32.2 on the host
833        // (2.8×) kept ALTERNATING through the whole diffusion stack, and
834        // because the alternation counter is shared per class in call
835        // order, one projection drew the CPU arm every single time — 9.9
836        // seconds a step on a kernel that needs 0.4. Both arms are
837        // compared on their BEST sample, so a 2× gap is not noise.
838        if (gn < PROBE_SAMPLES || cn < PROBE_SAMPLES) && g < cp * 2.0 && cp < g * 2.0 {
839            return;
840        }
841        let winner = if g <= cp { 1 } else { 2 };
842        if p.state
843            .compare_exchange(0, winner, Ordering::Relaxed, Ordering::Relaxed)
844            .is_ok()
845        {
846            tracing::info!(
847                "gpu probe [{}]: gpu {:.2} ms vs cpu {:.2} ms per op → {}",
848                class_name,
849                g / 1e6,
850                cp / 1e6,
851                if winner == 1 { "gpu" } else { "cpu" },
852            );
853            if let Some(c) = cache {
854                probe_cache_store(c, winner);
855            }
856        }
857    }
858}
859
860/// Is the class still collecting samples? (Call sites use this to route
861/// cold-weight calls away from the GPU arm during probing.)
862pub fn probe_deciding(c: OpClass) -> bool {
863    probe_on_for(c) && PROBES[c as usize].state.load(Ordering::Relaxed) == 0
864}
865
866/// Probing helper: true — tensor `idx`'s quant weights are ALREADY
867/// device-resident (a clean GPU sample is possible now); false — they
868/// were not (the upload starts within the VRAM budget, so a later call
869/// finds them warm) or the tensor cannot go to the GPU at all. Keeps the
870/// probe from billing a full cold dispatch+readback to a sample it will
871/// discard anyway. The verdict needs only a couple of warm tensors, so
872/// probe-driven uploads are capped — the losing-GPU machine should not
873/// pay for uploading the whole layer stack it will never use; if the GPU
874/// wins, the rest uploads lazily on demand, in the same first-touch order.
875#[allow(unused_variables)]
876pub fn q8_resident_or_upload(model: &Arc<CmfModel>, idx: usize) -> bool {
877    static PROBE_UPLOADS: AtomicU32 = AtomicU32::new(0);
878    let may_upload = PROBE_UPLOADS.load(Ordering::Relaxed) < 4;
879    let resident = match backend() {
880        #[cfg(target_os = "macos")]
881        Backend::Metal => crate::gpu_metal::q8_resident_or_upload(model, idx, may_upload),
882        #[cfg(feature = "gpu")]
883        Backend::Wgpu => crate::gpu_wgpu::q8_resident_or_upload(model, idx, may_upload),
884        Backend::None => false,
885    };
886    if !resident && may_upload {
887        PROBE_UPLOADS.fetch_add(1, Ordering::Relaxed);
888    }
889    resident
890}
891
892/// Test hook: reset all probes to the undecided state.
893#[cfg(test)]
894pub(crate) fn probe_reset() {
895    for p in &PROBES {
896        p.state.store(0, Ordering::Relaxed);
897        p.flip.store(0, Ordering::Relaxed);
898        p.gpu_ns.store(0, Ordering::Relaxed);
899        p.gpu_n.store(0, Ordering::Relaxed);
900        p.cpu_ns.store(0, Ordering::Relaxed);
901        p.cpu_n.store(0, Ordering::Relaxed);
902    }
903}
904
905/// The probe table is process-global by design, while these unit tests reset
906/// and seed selected entries to exercise arbitration. Keep only those tests
907/// out of each other's way; production callers still probe concurrently.
908#[cfg(test)]
909static PROBE_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
910
911#[cfg(test)]
912fn probe_test_guard() -> std::sync::MutexGuard<'static, ()> {
913    PROBE_TEST_LOCK
914        .lock()
915        .unwrap_or_else(std::sync::PoisonError::into_inner)
916}
917
918#[cfg(test)]
919mod probe_tests {
920    use super::*;
921    use std::time::Duration;
922
923    #[test]
924    fn cpu_only_whole_operator_dispatch_inherits_and_restores_scope() {
925        let pool = crate::pool::Pool::with_spin(3, 0);
926        cpu_scope(|| {
927            let inherit = inherit_cpu_scope();
928            pool.run_rows(64, &|_, _| {
929                let _guard = inherit();
930                assert!(CPU_ONLY.get());
931            });
932        });
933        pool.run_rows(64, &|_, _| assert!(!CPU_ONLY.get()));
934    }
935
936    // One test fn: PROBES is process-global and probe_reset touches all
937    // classes — parallel test threads would race.
938    #[test]
939    fn probe_alternates_discards_cold_and_decides() {
940        let _probe_guard = probe_test_guard();
941        probe_reset();
942        // Probing: arms alternate.
943        assert!(matches!(probe_arm(OpClass::Ffn), ProbeArm::Gpu));
944        assert!(matches!(probe_arm(OpClass::Ffn), ProbeArm::CpuTimed));
945
946        // A cold GPU sample (upload noted) must be discarded: feed a
947        // catastrophic cold sample, then clean fast-GPU samples — GPU
948        // wins only if the cold one did not count.
949        probe_note_cold();
950        probe_record(OpClass::Ffn, true, Duration::from_secs(1000));
951        for _ in 0..PROBE_SAMPLES {
952            probe_record(OpClass::Ffn, true, Duration::from_millis(1));
953            probe_record(OpClass::Ffn, false, Duration::from_millis(4));
954        }
955        assert!(matches!(probe_arm(OpClass::Ffn), ProbeArm::Gpu));
956
957        // The reverse: a class where the CPU arm is faster decides CPU.
958        for _ in 0..PROBE_SAMPLES {
959            probe_record(OpClass::Matmat, true, Duration::from_millis(4));
960            probe_record(OpClass::Matmat, false, Duration::from_millis(1));
961        }
962        assert!(matches!(probe_arm(OpClass::Matmat), ProbeArm::Cpu));
963
964        // cpu_scope: gates off inside, restored after.
965        cpu_scope(|| CPU_ONLY.with(|c| assert!(c.get())));
966        CPU_ONLY.with(|c| assert!(!c.get()));
967        cpu_scope(|| {
968            cpu_scope(|| CPU_ONLY.with(|c| assert!(c.get())));
969            CPU_ONLY.with(|c| assert!(c.get()));
970        });
971        let _ = std::panic::catch_unwind(|| cpu_scope(|| panic!("scope test")));
972        CPU_ONLY.with(|c| assert!(!c.get()));
973        probe_reset();
974    }
975
976    #[test]
977    fn a_remembered_verdict_is_adopted_and_a_stranger_is_not() {
978        let _probe_guard = probe_test_guard();
979        // Probing is not free: on a Snapdragon 778G the deciding classes
980        // cost minutes of wall clock before the first token, every
981        // process, and reached the same verdict every time. The cache
982        // exists so that price is paid once.
983        //
984        // The key is built from THIS process's device, never a name this
985        // test sets: `probe_set_device` is first-writer-wins and on a Mac
986        // the Metal backend may already have named the silicon before the
987        // tests run — which is exactly how this test failed on CI while
988        // passing locally. GemmNt on purpose: the arbitration test never
989        // touches it, and both run in one process.
990        let mine = probe_cache_key_named("gemm-nt");
991        let state = || {
992            PROBES[OpClass::GemmNt as usize]
993                .state
994                .load(Ordering::Relaxed)
995        };
996
997        // Another device's verdict is not mine, whatever it claims.
998        probe_cache_adopt("SomeOtherGPU/Vulkan\tgemm-nt\tgpu\n");
999        assert_eq!(state(), 0);
1000        // Neither is one from another build of this engine.
1001        let older = mine.replacen(env!("CARGO_PKG_VERSION"), "0.0.0-old", 1);
1002        assert_ne!(older, mine);
1003        probe_cache_adopt(&format!("{older}\tgpu\n"));
1004        assert_eq!(state(), 0);
1005        // Mine is.
1006        probe_cache_adopt(&format!("{mine}\tcpu\n"));
1007        assert_eq!(state(), 2);
1008
1009        PROBES[OpClass::GemmNt as usize]
1010            .state
1011            .store(0, Ordering::Relaxed);
1012    }
1013}
1014
1015/// Default row threshold: the GPU takes only larger matrices (lm_head
1016/// class). Below it, the dispatch/readback cost does not pay off on unified memory.
1017pub const GPU_MIN_ROWS: usize = 65_536;
1018
1019/// Effective threshold: `CMF_GPU_MIN_ROWS` overrides. Defaults differ
1020/// by device class: on a DISCRETE card VRAM bandwidth pays off even for
1021/// FFN/QKV-class matrices (4096), on unified memory only lm_head-class
1022/// is worth the dispatch/readback (65536). Field case behind this: a
1023/// 35B model on an RTX 4090 saw ~0 offload because every layer matrix
1024/// sat below the old universal 65536.
1025pub fn min_rows() -> usize {
1026    if let Some(v) = std::env::var("CMF_GPU_MIN_ROWS")
1027        .ok()
1028        .and_then(|v| v.parse().ok())
1029    {
1030        return v;
1031    }
1032    if discrete() { 4096 } else { GPU_MIN_ROWS }
1033}
1034
1035/// Is the active backend a discrete card (PCIe VRAM)?
1036pub fn discrete() -> bool {
1037    match backend() {
1038        #[cfg(feature = "gpu")]
1039        Backend::Wgpu => crate::gpu_wgpu::is_discrete(),
1040        #[cfg(target_os = "macos")]
1041        Backend::Metal => false, // UMA by the init() guard
1042        Backend::None => false,
1043    }
1044}
1045
1046/// A single MoE-FFN job (an expert with its own weight), executed in one
1047/// submission: (rows, cols, idx, row_scale) for gate/up/down + prescaled
1048/// inputs + the down column scale + the blending weight.
1049pub struct MoeJob<'a> {
1050    pub gate: (usize, usize, usize, &'a [f32]),
1051    pub up: (usize, usize, usize, &'a [f32]),
1052    pub down: (usize, usize, usize, &'a [f32]),
1053    pub xs_gate: Vec<f32>,
1054    pub xs_up: Vec<f32>,
1055    pub down_col: &'a [f32],
1056    pub w: f32,
1057    /// q1 trio: scales live inside the 6-byte tiles (row_scale slices
1058    /// empty, xs raw f32). Backends without a q1 kernel refuse the job.
1059    pub q1: bool,
1060    /// q4_tiled trio: scales inside the 18-byte tiles (row_scale
1061    /// slices empty, xs raw f32) — the MoE-hybrid coder class.
1062    pub q4t: bool,
1063    /// q4tp trio: same raw-xs contract, 16-byte nibble stride and the scale
1064    /// on a per-row ladder. Without this the experts of a q4tp MoE model fall
1065    /// to the CPU while every other dtype rides the device.
1066    pub q4tp: bool,
1067    /// Mixed 2-bit profile: gate/up are q2tp (8-byte chunks, zero rung),
1068    /// down stays q4tp. Set together with `q4tp`; a backend without the
1069    /// 2-bit kernel must refuse the whole job.
1070    pub gu_q2: bool,
1071    /// The reference's `swiglu_limit`; 0 disables the clamp. A backend that
1072    /// cannot apply it must REFUSE the job rather than drop it silently —
1073    /// the difference only shows on saturating activations, which is the
1074    /// hardest kind of divergence to notice.
1075    pub swiglu_limit: f32,
1076}
1077
1078/// A single independent batch matvec (GDN projections of one input).
1079pub struct BatchJob<'a> {
1080    pub idx: usize,
1081    pub rows: usize,
1082    pub cols: usize,
1083    pub row_scale: &'a [f32],
1084    pub xs: Vec<f32>,
1085    /// Weight layout. Was a bare `q1: bool`, which could only ever spell two
1086    /// of the four and silently sent everything else back to the CPU — the
1087    /// GDN projections of a q4t/q4tp model never reached the device at all.
1088    pub layout: BatchLayout,
1089}
1090
1091/// Which kernel a batched matvec needs. q8 carries row scales in a side
1092/// buffer; the rest embed them in the payload and differ in stride.
1093#[derive(Clone, Copy, PartialEq, Eq, Debug)]
1094pub enum BatchLayout {
1095    Q8,
1096    Q1,
1097    Q4t,
1098    Q4tp,
1099}
1100
1101#[derive(Clone, Copy, PartialEq, Eq)]
1102enum Backend {
1103    None,
1104    #[cfg(target_os = "macos")]
1105    Metal,
1106    #[cfg(feature = "gpu")]
1107    Wgpu,
1108}
1109
1110fn backend() -> Backend {
1111    #[cfg(feature = "gpu")]
1112    if crate::gpu_wgpu::selected() {
1113        return if crate::gpu_wgpu::enabled() {
1114            Backend::Wgpu
1115        } else {
1116            Backend::None
1117        };
1118    }
1119    #[cfg(target_os = "macos")]
1120    if crate::gpu_metal::enabled() {
1121        return Backend::Metal;
1122    }
1123    Backend::None
1124}
1125
1126/// GPU enabled and initialized on the selected backend?
1127/// Whether THIS build can bring a GPU up on THIS device: a compiled-in
1128/// backend plus a live adapter. The mobile FFI exposes it so an app can
1129/// tell "GPU off" from "GPU impossible" (a CPU-only .so ships no
1130/// backend at all). Cached after the first call.
1131pub fn backend_available() -> bool {
1132    #[cfg(target_os = "macos")]
1133    {
1134        // The Metal path is always compiled on macOS.
1135        true
1136    }
1137    #[cfg(all(feature = "gpu", not(target_os = "macos")))]
1138    {
1139        static AVAIL: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
1140        *AVAIL.get_or_init(crate::gpu_wgpu::adapter_probe)
1141    }
1142    #[cfg(all(not(feature = "gpu"), not(target_os = "macos")))]
1143    {
1144        false
1145    }
1146}
1147
1148/// A process-wide, phase-scoped GPU gate. `cpu_scope` is thread-local and
1149/// the pool's workers do not inherit it, so a caller that wants a whole
1150/// *phase* off the device — a prompt encoder whose weights live in a part of
1151/// the file the hot loop never touches, on a machine that cannot keep both
1152/// wired — has to say so globally.
1153static GPU_PAUSED: AtomicBool = AtomicBool::new(false);
1154
1155/// Park the device for every thread until the returned guard drops.
1156pub fn pause_gpu() -> GpuPause {
1157    GPU_PAUSED.store(true, Ordering::Relaxed);
1158    GpuPause(())
1159}
1160
1161pub struct GpuPause(());
1162
1163impl Drop for GpuPause {
1164    fn drop(&mut self) {
1165        GPU_PAUSED.store(false, Ordering::Relaxed);
1166    }
1167}
1168
1169pub fn enabled() -> bool {
1170    !GPU_PAUSED.load(Ordering::Relaxed) && backend() != Backend::None
1171}
1172
1173/// Default-on condition for the wgpu whole-token graph: the wgpu
1174/// backend on a DISCRETE adapter. NOT plain `enabled()` (macOS/Metal
1175/// must not pay a per-token layer scan for a graph its backend
1176/// refuses), and NOT integrated adapters: the graph's ~300 barriered
1177/// dispatches per token are cheap on desktop immediate-mode GPUs but
1178/// tiled mobile GPUs (Adreno/Mali) drain the pipeline at every barrier
1179/// — field report: 0.2 tok/s on-graph vs 15 tok/s on the CPU. On
1180/// integrated adapters the per-op probe path arbitrates each op class
1181/// against the CPU instead; CMF_GPU_WGPU_GRAPH=1 still forces the
1182/// graph anywhere.
1183/// Is the wgpu backend active at all (any adapter)? Eligibility gate
1184/// for the whole-token graph — whether it actually RUNS is decided by
1185/// `wgpu_graph_default` (trusted on discrete) or the generation race.
1186pub fn wgpu_active() -> bool {
1187    #[cfg(feature = "gpu")]
1188    {
1189        matches!(backend(), Backend::Wgpu)
1190    }
1191    #[cfg(not(feature = "gpu"))]
1192    {
1193        false
1194    }
1195}
1196
1197/// Which GPU this thread's engine calls address. Multi-card hosts hold
1198/// one wgpu context PER card (weights, KV mirrors and scratch live
1199/// inside a context, so per-device contexts give per-device caches for
1200/// free); this thread-local says which one is current. Default: the
1201/// process pin (CMF_GPU_ADAPTER) or 0 — so single-card runs behave
1202/// exactly as they always have.
1203pub fn default_device() -> usize {
1204    static D: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
1205    *D.get_or_init(|| {
1206        std::env::var("CMF_GPU_ADAPTER")
1207            .ok()
1208            .and_then(|v| v.trim().parse::<usize>().ok())
1209            .unwrap_or(0)
1210    })
1211}
1212
1213thread_local! {
1214    static CUR_DEV: std::cell::Cell<Option<usize>> = const { std::cell::Cell::new(None) };
1215}
1216
1217/// The device this thread is pinned to.
1218pub fn current_device() -> usize {
1219    CUR_DEV.with(|c| c.get()).unwrap_or_else(default_device)
1220}
1221
1222/// Pin this thread to a device. Server slots call it once per request;
1223/// the worker pool propagates it into its threads, so a dispatch begun
1224/// on card 1 does not finish on card 0.
1225pub fn set_current_device(i: usize) {
1226    CUR_DEV.with(|c| c.set(Some(i)));
1227}
1228
1229/// Run `f` with this thread pinned to `dev`, restoring the previous pin.
1230pub fn with_device<R>(dev: usize, f: impl FnOnce() -> R) -> R {
1231    let prev = CUR_DEV.with(|c| c.replace(Some(dev)));
1232    let r = f();
1233    CUR_DEV.with(|c| c.set(prev));
1234    r
1235}
1236
1237/// How many GPUs this process can address (wgpu adapter count; 1 on
1238/// Metal, 0 without a backend).
1239pub fn device_count() -> usize {
1240    #[cfg(all(feature = "gpu", not(target_os = "macos")))]
1241    {
1242        return crate::gpu_wgpu::adapter_count();
1243    }
1244    #[cfg(not(all(feature = "gpu", not(target_os = "macos"))))]
1245    {
1246        usize::from(backend_available())
1247    }
1248}
1249
1250/// Weight budget of the current GPU in bytes; 0 when there is none and
1251/// u64::MAX on unified memory (where the OS pages shared RAM and the
1252/// question "does the model fit the card" has no separate answer).
1253pub fn vram_budget() -> u64 {
1254    #[cfg(all(feature = "gpu", not(target_os = "macos")))]
1255    {
1256        return crate::gpu_wgpu::device_vram_budget();
1257    }
1258    #[cfg(not(all(feature = "gpu", not(target_os = "macos"))))]
1259    {
1260        if backend_available() { u64::MAX } else { 0 }
1261    }
1262}
1263
1264/// Bytes currently accounted as resident weight buffers on the active wgpu
1265/// adapter.  This is the logical device-local weight set; physical driver
1266/// allocations are reported separately by the platform tools.
1267pub fn resident_bytes() -> u64 {
1268    #[cfg(feature = "gpu")]
1269    {
1270        if backend() == Backend::Wgpu {
1271            return crate::gpu_wgpu::resident_bytes();
1272        }
1273    }
1274    0
1275}
1276
1277/// Sealed O(1) device mirror count and logical bytes for one pipeline id.
1278/// Zero is returned when wgpu is unavailable or the sequence has not reached
1279/// an O(1) seal yet.
1280pub fn o1_device_stats(kv_id: u64) -> (usize, u64) {
1281    #[cfg(feature = "gpu")]
1282    {
1283        if backend() == Backend::Wgpu {
1284            return crate::gpu_wgpu::o1_device_stats(kv_id);
1285        }
1286    }
1287    let _ = kv_id;
1288    (0, 0)
1289}
1290
1291/// Device weight bytes uploaded so far (wgpu; 0 on other backends).
1292/// Steady-state windows must show a ZERO delta — growth mid-benchmark
1293/// means eviction/re-upload and disqualifies the number.
1294pub fn upload_bytes() -> u64 {
1295    #[cfg(feature = "gpu")]
1296    {
1297        return crate::gpu_wgpu::UPLOAD_BYTES.load(std::sync::atomic::Ordering::Relaxed);
1298    }
1299    #[cfg(not(feature = "gpu"))]
1300    0
1301}
1302
1303/// Measure a transient host-to-device upload when the wgpu backend is
1304/// compiled in. CPU-only builds keep the benchmark command available and
1305/// report no device measurement instead of referring to the gated module.
1306pub fn upload_bandwidth_probe(block: usize, rounds: usize) -> Option<f64> {
1307    #[cfg(feature = "gpu")]
1308    {
1309        return crate::gpu_wgpu::upload_bandwidth_probe(block, rounds);
1310    }
1311    let _ = (block, rounds);
1312    None
1313}
1314
1315/// Which half of the run is asking.
1316///
1317/// The phase exists because the graph is plausibly two decisions, not
1318/// one — but on the hardware measured so far it is only ever a decode
1319/// decision. On an Adreno 642L with bonsai-1.7b, from identical clean
1320/// starts and two repeats each: decode 11.6 tok/s without it and 0.72
1321/// with, while prefill is 4.2 either way. A first reading of 3.4 -> 18.0
1322/// for prefill did not survive a controlled re-run — it was a dirty
1323/// probe cache between configurations, not the graph, and the prefill
1324/// route through the graph is GDN-only in the first place, which this
1325/// dense model never takes.
1326#[derive(Clone, Copy, PartialEq, Eq, Debug)]
1327pub enum GraphPhase {
1328    Prefill,
1329    Decode,
1330}
1331
1332/// The one place that decides whether the whole-token graph runs.
1333///
1334/// `CMF_GPU_WGPU_GRAPH`: `0` off everywhere, `prefill` only for the
1335/// prompt, anything else on everywhere. Unset: desktop-class GPUs take
1336/// it for both phases; phone-class UMA takes it for PREFILL only, which
1337/// is the measurement above rather than a guess — the per-op path keeps
1338/// decode, where it is seventeen times better.
1339pub fn wgpu_graph_on(phase: GraphPhase) -> bool {
1340    match std::env::var("CMF_GPU_WGPU_GRAPH").ok().as_deref() {
1341        Some("0") => false,
1342        Some("prefill") => phase == GraphPhase::Prefill,
1343        Some(_) => true,
1344        None => {
1345            if wgpu_graph_default() {
1346                return true;
1347            }
1348            // Integrated/mobile keeps the per-op path for BOTH phases —
1349            // unchanged, because the measurement that would have bought
1350            // prefill a graph did not reproduce. `=prefill` is there for
1351            // the device where it does; the default does not guess.
1352            let _ = phase;
1353            false
1354        }
1355    }
1356}
1357
1358pub fn wgpu_graph_default() -> bool {
1359    #[cfg(feature = "gpu")]
1360    {
1361        // Discrete cards always; Apple-silicon UMA on macOS too — desktop
1362        // -class GPUs where the graph measured ~2x the CPU on the Qwen3.6
1363        // family (M4: 13.3 tok/s against 7.3). Phone-class UMA (Android/
1364        // iOS builds) keeps the per-op probe path: tiled mobile GPUs have
1365        // turned the ~300-dispatch graph into seconds per token.
1366        matches!(backend(), Backend::Wgpu)
1367            && (crate::gpu_wgpu::discrete_active()
1368                || (cfg!(target_os = "macos") && crate::gpu_wgpu::adapter_up()))
1369    }
1370    #[cfg(not(feature = "gpu"))]
1371    {
1372        false
1373    }
1374}
1375
1376/// q8_row/q8_2f matvec, rows [row0, row0+rows). `xs` — prescaled by the column scale.
1377#[allow(clippy::too_many_arguments, unused_variables)]
1378pub fn q8_matvec_range(
1379    model: &Arc<CmfModel>,
1380    idx: usize,
1381    row0: usize,
1382    row_scale: &[f32],
1383    xs: &[f32],
1384    rows: usize,
1385    cols: usize,
1386    out: &mut [f32],
1387) -> bool {
1388    match backend() {
1389        #[cfg(target_os = "macos")]
1390        Backend::Metal => {
1391            crate::gpu_metal::q8_matvec_range(model, idx, row0, row_scale, xs, rows, cols, out)
1392        }
1393        #[cfg(feature = "gpu")]
1394        Backend::Wgpu => {
1395            crate::gpu_wgpu::q8_matvec_range(model, idx, row0, row_scale, xs, rows, cols, out)
1396        }
1397        Backend::None => false,
1398    }
1399}
1400
1401/// Decode-exact short q8_2f panel for the banked MiMo head.
1402pub(crate) fn q82_short_rows(
1403    model: &Arc<CmfModel>,
1404    idx: usize,
1405    xs: &[f32],
1406    b: usize,
1407    rows: usize,
1408    cols: usize,
1409    out: &mut [f32],
1410) -> bool {
1411    #[cfg(feature = "gpu")]
1412    if enabled_here() && backend() == Backend::Wgpu {
1413        return crate::gpu_wgpu::q82_short_rows(model, idx, xs, b, rows, cols, out);
1414    }
1415    let _ = (model, idx, xs, b, rows, cols, out);
1416    false
1417}
1418
1419/// GEMM of a prefill batch: `pre` — prescaled inputs row-major [b, cols],
1420/// out — row-major [b, rows].
1421#[allow(clippy::too_many_arguments, unused_variables)]
1422/// The two-field int8 GEMM with the column field left for the device.
1423/// wgpu only — Metal's int8 kernel takes a pre-scaled activation, so the
1424/// caller keeps that path when this returns `false`.
1425#[allow(clippy::too_many_arguments)]
1426pub fn q8_matmat_2f(
1427    model: &Arc<CmfModel>,
1428    idx: usize,
1429    row_scale: &[f32],
1430    col_field: &[f32],
1431    xs: &[f32],
1432    b: usize,
1433    rows: usize,
1434    cols: usize,
1435    out: &mut [f32],
1436) -> bool {
1437    #[allow(unreachable_patterns)]
1438    match backend() {
1439        #[cfg(feature = "gpu")]
1440        Backend::Wgpu => {
1441            crate::gpu_wgpu::q8_matmat_2f(model, idx, row_scale, col_field, xs, b, rows, cols, out)
1442        }
1443        _ => false,
1444    }
1445}
1446
1447pub fn q8_matmat(
1448    model: &Arc<CmfModel>,
1449    idx: usize,
1450    row_scale: &[f32],
1451    pre: &[f32],
1452    b: usize,
1453    rows: usize,
1454    cols: usize,
1455    out: &mut [f32],
1456) -> bool {
1457    match backend() {
1458        #[cfg(target_os = "macos")]
1459        Backend::Metal => {
1460            crate::gpu_metal::q8_matmat(model, idx, row_scale, pre, b, rows, cols, out)
1461        }
1462        #[cfg(feature = "gpu")]
1463        Backend::Wgpu => crate::gpu_wgpu::q8_matmat(model, idx, row_scale, pre, b, rows, cols, out),
1464        Backend::None => false,
1465    }
1466}
1467
1468/// q1 matvec: raw f32 activations, tile-embedded scales. Metal only
1469/// for now (wgpu q1 WGSL is queued); false = CPU fallback.
1470#[allow(unused_variables)]
1471pub fn q1_matvec(
1472    model: &Arc<CmfModel>,
1473    idx: usize,
1474    xs: &[f32],
1475    rows: usize,
1476    cols: usize,
1477    out: &mut [f32],
1478) -> bool {
1479    match backend() {
1480        #[cfg(target_os = "macos")]
1481        Backend::Metal => crate::gpu_metal::q1_matvec(model, idx, xs, rows, cols, out),
1482        #[cfg(feature = "gpu")]
1483        Backend::Wgpu => crate::gpu_wgpu::q1_matvec(model, idx, xs, rows, cols, out),
1484        Backend::None => false,
1485    }
1486}
1487
1488/// Whole attention sub-block on the wgpu token graph (drop-in for
1489/// `qwen_attention`): normed hidden in, O-projection out, resident device
1490/// K/V mirror. false = refusal / not the wgpu backend → CPU path.
1491#[allow(clippy::too_many_arguments)]
1492pub fn attn_dropin(
1493    model: &Arc<CmfModel>,
1494    kv_id: u64,
1495    layer: usize,
1496    normed: &[f32],
1497    wq_idx: usize,
1498    wk_idx: usize,
1499    wv_idx: usize,
1500    wo_idx: usize,
1501    q_norm: Option<&[f32]>,
1502    k_norm: Option<&[f32]>,
1503    late_qk_norm: bool,
1504    invf: &[f32],
1505    nh: usize,
1506    nkv: usize,
1507    hd: usize,
1508    rd: usize,
1509    hidden: usize,
1510    pos: usize,
1511    cap: usize,
1512    gemma: bool,
1513    eps: f32,
1514    cpu_k: &[Vec<f32>],
1515    cpu_v: &[Vec<f32>],
1516    out: &mut [f32],
1517) -> bool {
1518    match backend() {
1519        #[cfg(feature = "gpu")]
1520        Backend::Wgpu => crate::gpu_wgpu::attn_dropin_gpu(
1521            model, kv_id, layer, normed, wq_idx, wk_idx, wv_idx, wo_idx, q_norm, k_norm,
1522            late_qk_norm, invf, nh, nkv, hd, rd, hidden, pos, cap, gemma, eps, cpu_k, cpu_v, out,
1523        ),
1524        #[allow(unused_variables)]
1525        _ => false,
1526    }
1527}
1528
1529/// Descriptor operation attached to a graph weight.  `None` is the default
1530/// for ordinary CMF files; Prism weights are admitted only when the token
1531/// graph carries this explicit transform contract.
1532#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1533pub enum GraphPrismOp {
1534    None,
1535    Forward,
1536    InverseEmbedding,
1537}
1538
1539/// One weight in the whole-token graph: tensor idx + a codec tag (0=q8_row,
1540/// 1=q1, 2=q4_tiled, 3=q1t, 4=f32) + per-row scales (q8_row only) + the raw f32
1541/// data (kind 4 only — small unquantized projections like GDN in_proj_a/b).
1542pub struct GraphW<'a> {
1543    pub idx: usize,
1544    pub kind: u8,
1545    pub row_scale: &'a [f32],
1546    pub data: &'a [f32],
1547    pub prism: GraphPrismOp,
1548    pub affine: bool,
1549}
1550
1551/// A layer's token-mixing op: standard attention or a GDN (linear-attention)
1552/// block. The surrounding norms + SwiGLU FFN are common to both.
1553pub enum GraphAttn<'a> {
1554    Full {
1555        wq: GraphW<'a>,
1556        wk: GraphW<'a>,
1557        wv: GraphW<'a>,
1558        wo: GraphW<'a>,
1559        q_norm: Option<&'a [f32]>,
1560        k_norm: Option<&'a [f32]>,
1561        /// HunYuan dense: q/k norm after RoPE (rope-kernel flag bit 32).
1562        late_qk_norm: bool,
1563        /// (bq, bk, bv) attention biases (Qwen2). None ⇒ no bias.
1564        bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
1565        /// Qwen3.5 gated attention: wq emits 2·nh·hd (q||gate per head), the
1566        /// attention output is scaled by sigmoid(gate) before the O projection.
1567        output_gate: bool,
1568        cpu_k: &'a [Vec<f32>],
1569        cpu_v: &'a [Vec<f32>],
1570        /// This layer's own attention geometry, when the model's layers do
1571        /// not share one (MiMo-V2: 4/8 KV heads, 128-wide V under 192-wide
1572        /// heads, sliding windows with learned sinks, two RoPE tables).
1573        /// None = the call-wide (nkv, hd, rd, invf), V as wide as K, full
1574        /// context and a plain softmax — the historical contract, whose
1575        /// kernels and dispatch are untouched.
1576        geom: Option<GraphAttnGeom<'a>>,
1577    },
1578    Gdn {
1579        qkv: GraphW<'a>,
1580        z: GraphW<'a>,
1581        a: GraphW<'a>,
1582        b: GraphW<'a>,
1583        out: GraphW<'a>,
1584        conv1d: &'a [f32],
1585        a_log: &'a [f32],
1586        dt_bias: &'a [f32],
1587        norm: &'a [f32],
1588        nv: usize,
1589        nk: usize,
1590        dk: usize,
1591        dv: usize,
1592        kk: usize,
1593        /// CPU recurrent state `[ring (kk-1)·cdim | S nv·dk·dv]` — seeds the
1594        /// device mirror when prefill ran on the host (o1 collection, CPU
1595        /// fallback): a zero-initialized device state at decode is exactly
1596        /// the "coherent but contextless" garble.
1597        cpu_state: &'a [f32],
1598    },
1599    /// LFM2 gated short convolution: a fused (B, C, x) projection, a
1600    /// depthwise causal conv over a (kernel−1)-deep per-channel ring,
1601    /// C-gating, and an output projection. This mixer is what most of an
1602    /// LFM2 stack is (22 of the 2.6B's 30 layers), and before it had a
1603    /// graph arm the whole model fell to the per-op path — ~100 submits
1604    /// a token, 22 tok/s on an A100 for a 1.4 GB file.
1605    ShortConv {
1606        /// [3·hidden, hidden] fused input projection.
1607        inp: GraphW<'a>,
1608        /// [hidden, hidden] output projection.
1609        out: GraphW<'a>,
1610        /// [hidden · kernel] depthwise taps, `[channel][tap]`, tap
1611        /// kernel−1 multiplying the current position.
1612        taps: &'a [f32],
1613        kernel: usize,
1614        /// CPU conv ring `[channel][kernel−1]`, slot 0 newest — seeds
1615        /// the device mirror when prefill ran on the host, which for
1616        /// this mixer is always (the batch graph declines it).
1617        cpu_state: &'a [f32],
1618    },
1619}
1620
1621/// Per-layer attention geometry for the wgpu graphs (see
1622/// `GraphAttn::Full::geom`). The layer's CPU cache keeps K rows `hd` wide
1623/// and V rows zero-padded to `hd`; the device mirror stores V `dv` wide,
1624/// and a windowed layer keeps a ring of the last positions only.
1625#[derive(Clone, Copy)]
1626pub struct GraphAttnGeom<'a> {
1627    /// KV heads of this layer (divides the Q heads).
1628    pub nkv: usize,
1629    /// V head width, `4 <= dv <= head_dim`, a multiple of 4.
1630    pub dv: usize,
1631    /// Rotary width of this layer (NeoX half-split over `[0, rd)`).
1632    pub rd: usize,
1633    /// This layer's RoPE inverse frequencies (`rd / 2` of them).
1634    pub invf: &'a [f32],
1635    /// Positions a query sees, its own included (MiMo-V2 SWA: 128);
1636    /// None = the whole context.
1637    pub window: Option<usize>,
1638    /// Learned per-Q-head sink logits (gpt-oss / MiMo-V2): they join the
1639    /// softmax max and denominator and carry no value row.
1640    pub sink: Option<&'a [f32]>,
1641}
1642
1643/// Per-layer weights for the whole-token wgpu graph.
1644pub struct GraphLayer<'a> {
1645    pub input_norm: &'a [f32],
1646    pub attn: GraphAttn<'a>,
1647    pub post_norm: &'a [f32],
1648    pub ffn: GraphFfn<'a>,
1649}
1650
1651/// The FFN of one graph layer: a dense SwiGLU trio, or a routed MoE —
1652/// router + top-k selection + all selected experts run ON DEVICE (the
1653/// routing decision depends on the resident hidden state, so a CPU
1654/// round-trip per layer would forfeit the one-submit design).
1655pub enum GraphFfn<'a> {
1656    /// A singleton attention-only batch graph. Returns the post-attention
1657    /// residual, allowing a dynamic expert bank to own the FFN separately.
1658    AttentionOnly,
1659    Dense {
1660        gate: GraphW<'a>,
1661        up: GraphW<'a>,
1662        down: GraphW<'a>,
1663    },
1664    Moe {
1665        /// Router logits weight (f32, kind 4) `[n_exp, hidden]`.
1666        router: GraphW<'a>,
1667        /// Shared-expert sigmoid gate (f32) `[1, hidden]`.
1668        shared_gate: GraphW<'a>,
1669        /// Per-expert q4_tiled directory indices `(gate, up, down)`;
1670        /// the SHARED expert rides as the LAST entry — the select
1671        /// kernel pins it with the sigmoid weight.
1672        experts: Vec<(usize, usize, usize)>,
1673        /// Routed experts (shared excluded).
1674        n_exp: usize,
1675        top_k: usize,
1676        inter: usize,
1677        norm_topk: bool,
1678        /// Expert weight layout, uniform across the layer: `false` =
1679        /// q4_tiled (18 B tiles, inline f16 scale), `true` = q4tp
1680        /// (16 B nibbles + a per-row ladder plane). The two differ only
1681        /// in where the scale comes from, so they share every kernel
1682        /// but the weight-staging block.
1683        q4tp: bool,
1684        /// `true` = the gate/up experts are `q2tp` (2-bit plane) while
1685        /// `down` stays q4tp — the mixed profile a 2-bit-class checkpoint
1686        /// converts into. Only meaningful with `q4tp: true`.
1687        gu_q2: bool,
1688        /// LFM2-MoE / DeepSeek-V3 `noaux_tc` routing: per-expert sigmoid
1689        /// scores instead of a softmax, and `norm_topk` renormalises with
1690        /// the 1e-6 floor. The softmax arm is bit-identical to before.
1691        sigmoid: bool,
1692        /// Per-expert SELECTION bias: added to the score for the top-k
1693        /// choice only — the mixing weights stay unbiased (noaux_tc).
1694        bias: Option<&'a [f32]>,
1695        /// Whether a shared expert rides as the last `experts` entry.
1696        /// LFM2-MoE has none; the select kernel then leaves slot `top_k`
1697        /// unwritten and the expert loop runs `top_k` slots, not +1.
1698        has_shared: bool,
1699        /// The shared expert carries a sigmoid gate (Qwen2/3-MoE). `false`
1700        /// with `has_shared`: the shared expert enters with weight 1
1701        /// (DeepSeek-V3 / HunYuan hy_v3) and `shared_gate` is a stand-in
1702        /// the select kernels ignore.
1703        shared_gated: bool,
1704        /// Multiplier on the routed mixing weights after the optional
1705        /// renormalization (`routed_scaling_factor`); 1.0 = none.
1706        route_scale: f32,
1707    },
1708}
1709
1710/// Outcome of one whole-token graph attempt. A failed attempt after sealed
1711/// O(1) state was admitted must not fall through to the stale CPU state.
1712#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1713pub enum TokenGraphOutcome {
1714    /// No command was committed; the caller may use its ordinary path.
1715    Declined,
1716    /// The graph completed and its hidden/logits output is valid.
1717    Completed,
1718    /// Sealed O(1) state was admitted and a later graph operation failed.
1719    Failed,
1720}
1721
1722/// Whole-token decode graph on wgpu: the entire layer stack in ONE submit,
1723/// hidden resident, one readback. Updates `h` in place.
1724/// `loop_norm_at`: virtual layer indices after which `final_norm` is applied
1725/// (Looped Transformer mid-stack norm). Empty for standard models.
1726#[allow(clippy::too_many_arguments)]
1727pub fn forward_token_graph(
1728    model: &Arc<CmfModel>,
1729    kv_id: u64,
1730    layers: &[GraphLayer],
1731    // Per-layer sealed o1 (Nystrom) state; Some = replace this layer's
1732    // exact attention with the O(1) kernels. wgpu only.
1733    o1: &[Option<Vec<crate::nystrom::O1DeviceView<'_>>>],
1734    o1_epoch: u64,
1735    invf: &[f32],
1736    h: &mut [f32],
1737    nh: usize,
1738    nkv: usize,
1739    hd: usize,
1740    attn_scale: f32,
1741    rd: usize,
1742    hidden: usize,
1743    inter: usize,
1744    position: usize,
1745    cap: usize,
1746    gemma: bool,
1747    eps: f32,
1748    lm_head: Option<(&GraphW, usize)>,
1749    final_norm: &[f32],
1750    logits: &mut Vec<f32>,
1751    loop_norm_at: &[usize],
1752    steps: usize,
1753    embed: Option<(&GraphW, usize, f32)>,
1754    ids_out: Option<&mut Vec<u32>>,
1755    // How many leading layers the graph ran (see the wgpu twin) — smaller
1756    // than layers.len() when the expert budget ended the device prefix.
1757    layers_run: Option<&mut usize>,
1758    // Absolute index of layers[0] in the model — the KV/state mirrors key
1759    // on it, so a layer SPAN (network split segment) shares mirrors with
1760    // a full-stack run instead of colliding at slot 0.
1761    layer_base: usize,
1762    // Read the final hidden back alongside the fused head's logits.
1763    hidden_too: bool,
1764) -> TokenGraphOutcome {
1765    match backend() {
1766        #[cfg(feature = "gpu")]
1767        Backend::Wgpu => crate::gpu_wgpu::forward_token_graph(
1768            model,
1769            kv_id,
1770            layers,
1771            o1,
1772            o1_epoch,
1773            invf,
1774            h,
1775            nh,
1776            nkv,
1777            hd,
1778            attn_scale,
1779            rd,
1780            hidden,
1781            inter,
1782            position,
1783            cap,
1784            gemma,
1785            eps,
1786            lm_head,
1787            final_norm,
1788            logits,
1789            loop_norm_at,
1790            steps,
1791            embed,
1792            ids_out,
1793            layers_run,
1794            layer_base,
1795            hidden_too,
1796        ),
1797        #[allow(unused_variables)]
1798        _ => {
1799            let _ = (
1800                attn_scale,
1801                lm_head,
1802                final_norm,
1803                logits,
1804                loop_norm_at,
1805                layers_run,
1806                layer_base,
1807                hidden_too,
1808            );
1809            TokenGraphOutcome::Declined
1810        }
1811    }
1812}
1813
1814/// Whole-token resident graph for the Embryo working model.  Unlike the
1815/// generic graph this path accepts a packed f32 model and keeps both phase
1816/// recurrent state and the anchor KV cache on the device.  `false` is an
1817/// honest capability refusal; callers must retain the CPU executor.
1818pub fn forward_embryo_graph(
1819    model: &Arc<EmbryoGraphModel>,
1820    kv_id: u64,
1821    hidden: &[f32],
1822    position: usize,
1823    logits: &mut Vec<f32>,
1824) -> bool {
1825    #[cfg(feature = "gpu")]
1826    if backend() == Backend::Wgpu {
1827        return crate::gpu_wgpu::forward_embryo_graph(model, kv_id, hidden, position, logits);
1828    }
1829    let _ = (model, kv_id, hidden, position, logits);
1830    false
1831}
1832
1833/// Speculative-verify tail for the batched graph: fold final-norm + lm_head
1834/// over every batch position and read all k logit rows back; the batch also
1835/// snapshots the GDN state per position for `gdn_spec_restore`.
1836#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1837pub enum BatchGraphOutcome {
1838    /// The graph declined before mutating persistent device state. Callers may
1839    /// safely use the existing per-position path.
1840    Declined,
1841    /// The complete batch committed and its readback succeeded.
1842    Completed,
1843    /// A batch that had admitted sealed O(1) state failed after admission.
1844    /// Falling back to CPU would mix two state machines, so the caller must
1845    /// abort and clear the sequence instead.
1846    Failed,
1847}
1848
1849pub struct SpecTail<'a> {
1850    pub lm: GraphW<'a>,
1851    pub lm_rows: usize,
1852    pub final_norm: &'a [f32],
1853    pub logits_out: &'a mut Vec<f32>,
1854}
1855
1856/// Batched prefill: k contiguous positions through the whole graph in one submit
1857/// (projections/FFN as GEMMs, attention/GDN looped over scratch). `h` is
1858/// [k·hidden] in/out; `positions` len k. wgpu only.
1859#[allow(clippy::too_many_arguments)]
1860pub fn forward_batch_graph(
1861    model: &Arc<CmfModel>,
1862    kv_id: u64,
1863    layers: &[GraphLayer],
1864    invf: &[f32],
1865    h: &mut [f32],
1866    nh: usize,
1867    nkv: usize,
1868    hd: usize,
1869    rd: usize,
1870    hidden: usize,
1871    inter: usize,
1872    positions: &[usize],
1873    cap: usize,
1874    gemma: bool,
1875    eps: f32,
1876    attn_scale: f32,
1877    k: usize,
1878    // Per-layer sealed O(1) device views. An empty slice means the ordinary
1879    // exact-KV path; otherwise it must have one entry per graph layer.
1880    o1: &[Option<Vec<crate::nystrom::O1DeviceView<'_>>>],
1881    o1_epoch: u64,
1882    spec: Option<SpecTail<'_>>,
1883    // Device-prefix mode (plain prefill only): Some = when the whole stack
1884    // does not fit the weight budget, run the leading layers that do — the
1885    // same prefix rule as the token graph — leave the boundary hidden in
1886    // `h` and report the count here; the caller runs the rest on the host.
1887    // None = all layers or a decline, as before.
1888    layers_run: Option<&mut usize>,
1889) -> BatchGraphOutcome {
1890    forward_batch_graph_at(model, kv_id, 0, layers, invf, h, nh, nkv, hd, rd, hidden, inter, positions, cap, gemma, eps, attn_scale, k, o1, o1_epoch, spec, layers_run)
1891}
1892
1893thread_local! {
1894    static MIMO_ATTN_SCRATCH: std::cell::Cell<bool> = const { std::cell::Cell::new(true) };
1895}
1896
1897pub(crate) fn mimo_attention_scratch_enabled() -> bool {
1898    MIMO_ATTN_SCRATCH.with(std::cell::Cell::get)
1899}
1900
1901/// Scoped diagnostic A/B switch; unlike process environment mutations it
1902/// cannot race a model's background workers. Restores on panic as well.
1903#[doc(hidden)]
1904pub fn mimo_attention_scratch_scope<R>(enabled: bool, f: impl FnOnce() -> R) -> R {
1905    struct Restore(bool);
1906    impl Drop for Restore {
1907        fn drop(&mut self) {
1908            MIMO_ATTN_SCRATCH.with(|v| v.set(self.0));
1909        }
1910    }
1911    let _restore = Restore(MIMO_ATTN_SCRATCH.with(|v| v.replace(enabled)));
1912    f()
1913}
1914
1915thread_local! {
1916    static MIMO_Q8_SHORT: std::cell::Cell<bool> = const { std::cell::Cell::new(true) };
1917}
1918
1919pub(crate) fn mimo_q8_short_enabled() -> bool {
1920    MIMO_Q8_SHORT.with(std::cell::Cell::get)
1921}
1922
1923/// Diagnostic A/B switch for row-exact short q8 graph kernels.
1924#[doc(hidden)]
1925pub fn mimo_q8_short_scope<R>(enabled: bool, f: impl FnOnce() -> R) -> R {
1926    struct Restore(bool);
1927    impl Drop for Restore {
1928        fn drop(&mut self) {
1929            MIMO_Q8_SHORT.with(|v| v.set(self.0));
1930        }
1931    }
1932    let _restore = Restore(MIMO_Q8_SHORT.with(|v| v.replace(enabled)));
1933    f()
1934}
1935
1936/// Batched graph over a span whose first absolute layer is `layer_base`.
1937#[allow(clippy::too_many_arguments)]
1938pub fn forward_batch_graph_at(
1939    model: &Arc<CmfModel>,
1940    kv_id: u64,
1941    layer_base: usize,
1942    layers: &[GraphLayer],
1943    invf: &[f32],
1944    h: &mut [f32],
1945    nh: usize,
1946    nkv: usize,
1947    hd: usize,
1948    rd: usize,
1949    hidden: usize,
1950    inter: usize,
1951    positions: &[usize],
1952    cap: usize,
1953    gemma: bool,
1954    eps: f32,
1955    attn_scale: f32,
1956    k: usize,
1957    // Per-layer sealed O(1) device views. An empty slice means the ordinary
1958    // exact-KV path; otherwise it must have one entry per graph layer.
1959    o1: &[Option<Vec<crate::nystrom::O1DeviceView<'_>>>],
1960    o1_epoch: u64,
1961    spec: Option<SpecTail<'_>>,
1962    // Device-prefix mode (plain prefill only): Some = when the whole stack
1963    // does not fit the weight budget, run the leading layers that do — the
1964    // same prefix rule as the token graph — leave the boundary hidden in
1965    // `h` and report the count here; the caller runs the rest on the host.
1966    // None = all layers or a decline, as before.
1967    layers_run: Option<&mut usize>,
1968) -> BatchGraphOutcome {
1969    match backend() {
1970        #[cfg(feature = "gpu")]
1971        Backend::Wgpu => crate::gpu_wgpu::forward_batch_graph_at(
1972            model, kv_id, layer_base, layers, invf, h, nh, nkv, hd, rd, hidden, inter, positions,
1973            cap, gemma, eps, attn_scale, k, o1, o1_epoch, spec, layers_run,
1974        ),
1975        #[allow(unreachable_patterns)]
1976        _ => {
1977            let _ = (o1, o1_epoch, spec, layers_run);
1978            BatchGraphOutcome::Declined
1979        }
1980    }
1981}
1982
1983/// After a partial speculative acceptance: restore every GDN layer's device
1984/// state to the snapshot after batch position `slot`. `base_pos` is the
1985/// absolute position of the first verify row and `expected_layers` makes the
1986/// restore all-or-nothing across the model's recurrent layers. wgpu only.
1987pub fn gdn_spec_restore(kv_id: u64, slot: usize, base_pos: usize, expected_layers: usize) -> bool {
1988    #[cfg(feature = "gpu")]
1989    if backend() == Backend::Wgpu {
1990        return crate::gpu_wgpu::gdn_spec_restore(kv_id, slot, base_pos, expected_layers);
1991    }
1992    #[allow(unreachable_code)]
1993    {
1994        let _ = (kv_id, slot, base_pos, expected_layers);
1995        false
1996    }
1997}
1998
1999/// Re-point one exact-attention device mirror after a speculative round has
2000/// discarded unaccepted rows. The rows beyond `stored` remain allocated and
2001/// are overwritten by the next append; only the logical cursor moves. This
2002/// is the wgpu twin of Metal's existing mirror cursor helper and keeps the
2003/// MTP graph's speculative/device cache coherent with its real anchor.
2004pub fn graph_kv_set_stored(kv_id: u64, layer: usize, stored: usize) -> bool {
2005    #[cfg(feature = "gpu")]
2006    if backend() == Backend::Wgpu {
2007        return crate::gpu_wgpu::kv_mirror_set_stored(kv_id, layer, stored);
2008    }
2009    #[cfg(target_os = "macos")]
2010    if backend() == Backend::Metal {
2011        crate::gpu_metal::kv_mirror_set_stored(kv_id, layer, stored);
2012        return true;
2013    }
2014    false
2015}
2016
2017/// Rows the wgpu token graph's exact-attention mirror holds for one layer
2018/// (None: no wgpu mirror). Metal keeps its owner cache current per token
2019/// and reports None here.
2020pub fn graph_kv_stored(_kv_id: u64, _layer: usize) -> Option<usize> {
2021    #[cfg(feature = "gpu")]
2022    if backend() == Backend::Wgpu {
2023        return crate::gpu_wgpu::kv_mirror_stored(_kv_id, _layer);
2024    }
2025    None
2026}
2027
2028/// Does the wgpu token graph hold a device-resident recurrent state for
2029/// this layer (one the host `linear_state` has not seen)?
2030pub fn graph_state_resident(_kv_id: u64, _layer: usize) -> bool {
2031    #[cfg(feature = "gpu")]
2032    if backend() == Backend::Wgpu {
2033        return crate::gpu_wgpu::graph_state_resident(_kv_id, _layer);
2034    }
2035    false
2036}
2037
2038/// Copy rows back from the wgpu token graph's K/V mirrors in one submit:
2039/// for each `(layer, from, to)` the K and V rows `[from..to)`, position-major
2040/// (`[(to − from) × nkv × hd]` each).
2041pub fn graph_kv_read_rows(
2042    _kv_id: u64,
2043    _reqs: &[(usize, usize, usize)],
2044    _nkv: usize,
2045    _hd: usize,
2046) -> Option<Vec<(Vec<f32>, Vec<f32>)>> {
2047    #[cfg(feature = "gpu")]
2048    if backend() == Backend::Wgpu {
2049        return crate::gpu_wgpu::kv_mirror_read_rows(_kv_id, _reqs, _nkv, _hd);
2050    }
2051    None
2052}
2053
2054/// Rows `[from, to)` of one wgpu exact-attention mirror in the host
2055/// cache's layout (V zero-padded to `hd`), whatever the mirror's geometry
2056/// (narrow V, a sliding layer's ring). The third value is the first
2057/// position actually read: a ring returns zeros below it. None: no wgpu
2058/// mirror holding those rows at (nkv, hd).
2059pub fn graph_kv_pull_host(
2060    _kv_id: u64,
2061    _layer: usize,
2062    _from: usize,
2063    _to: usize,
2064    _nkv: usize,
2065    _hd: usize,
2066) -> Option<(Vec<f32>, Vec<f32>, usize)> {
2067    #[cfg(feature = "gpu")]
2068    if backend() == Backend::Wgpu {
2069        return crate::gpu_wgpu::kv_mirror_pull_host(_kv_id, _layer, _from, _to, _nkv, _hd);
2070    }
2071    None
2072}
2073
2074/// Drop the wgpu token graph's device K/V mirror for a pipeline.
2075pub fn graph_kv_reset(_kv_id: u64) {
2076    #[cfg(feature = "gpu")]
2077    if backend() == Backend::Wgpu {
2078        crate::gpu_wgpu::kv_mirror_reset(_kv_id);
2079        crate::gpu_wgpu::embryo_graph_reset(_kv_id);
2080    }
2081}
2082
2083/// Ternary (q1t) BASE matvec on the GPU — fills `out` with the base dot; the
2084/// caller adds the sparse overlay on the CPU. Metal only for now (wgpu q1t not
2085/// yet written → CPU fallback).
2086pub fn q1t_matvec(
2087    model: &Arc<CmfModel>,
2088    idx: usize,
2089    xs: &[f32],
2090    rows: usize,
2091    cols: usize,
2092    out: &mut [f32],
2093) -> bool {
2094    match backend() {
2095        #[cfg(target_os = "macos")]
2096        Backend::Metal => {
2097            if metal_q1t_enabled() {
2098                crate::gpu_metal::q1t_matvec(model, idx, xs, rows, cols, out)
2099            } else {
2100                false
2101            }
2102        }
2103        #[cfg(feature = "gpu")]
2104        Backend::Wgpu => crate::gpu_wgpu::q1t_matvec(model, idx, xs, rows, cols, out),
2105        Backend::None => false,
2106    }
2107}
2108
2109/// q4_block matvec on the GPU — wgpu only (Metal drives q4_block through the
2110/// whole-token graph, not a standalone matvec).
2111#[allow(unused_variables)]
2112pub fn q4b_matvec(
2113    model: &Arc<CmfModel>,
2114    idx: usize,
2115    xs: &[f32],
2116    rows: usize,
2117    cols: usize,
2118    out: &mut [f32],
2119) -> bool {
2120    match backend() {
2121        #[cfg(target_os = "macos")]
2122        Backend::Metal => false,
2123        #[cfg(feature = "gpu")]
2124        Backend::Wgpu => crate::gpu_wgpu::q4b_matvec(model, idx, xs, rows, cols, out),
2125        Backend::None => false,
2126    }
2127}
2128
2129/// q1t batched GEMM (prefill) — base + overlay on-device (Metal simdgroup or
2130/// wgpu register-blocked).
2131pub fn q1t_matmat(
2132    model: &Arc<CmfModel>,
2133    idx: usize,
2134    xs: &[f32],
2135    b: usize,
2136    rows: usize,
2137    cols: usize,
2138    out: &mut [f32],
2139) -> bool {
2140    match backend() {
2141        #[cfg(target_os = "macos")]
2142        // Batched prefill and single-token decode are both enabled. On the
2143        // real 14.8B Q1T model prefill PPL was within 0.3% of CPU (7.942 vs
2144        // 7.966), and the alignment-safe decode kernel reached 3.52e-6 max_rel.
2145        Backend::Metal => crate::gpu_metal::q1t_matmat(model, idx, xs, b, rows, cols, out),
2146        #[cfg(feature = "gpu")]
2147        Backend::Wgpu => crate::gpu_wgpu::q1t_matmat(model, idx, xs, b, rows, cols, out),
2148        Backend::None => false,
2149    }
2150}
2151
2152/// Native Metal Q1T switch. Enabled by default after the byte-packed Q1T
2153/// fields were changed to alignment-safe loads; keep an explicit emergency
2154/// fallback for device/driver diagnostics.
2155#[cfg(target_os = "macos")]
2156pub(crate) fn metal_q1t_enabled() -> bool {
2157    std::env::var("CMF_METAL_Q1T")
2158        .map(|v| v != "0" && !v.eq_ignore_ascii_case("off"))
2159        .unwrap_or(true)
2160}
2161
2162/// Batched q1 GEMM (prefill). wgpu only — Metal has its own block path.
2163pub fn q1_matmat(
2164    model: &Arc<CmfModel>,
2165    idx: usize,
2166    xs: &[f32],
2167    b: usize,
2168    rows: usize,
2169    cols: usize,
2170    out: &mut [f32],
2171) -> bool {
2172    match backend() {
2173        #[cfg(feature = "gpu")]
2174        Backend::Wgpu => crate::gpu_wgpu::q1_matmat(model, idx, xs, b, rows, cols, out),
2175        #[allow(unused_variables)]
2176        _ => false,
2177    }
2178}
2179
2180/// Contention kill for the wide imagegen GEMM/FFN paths: one grossly
2181/// slow op under a work-proportional budget (fair-device ops are
2182/// ≤~100 ms even at 1024px) means another process owns the device —
2183/// verdicts are per-process, so CPU for the rest of this one.
2184static MM_KILL: AtomicBool = AtomicBool::new(false);
2185pub(crate) fn mm_killed() -> bool {
2186    MM_KILL.load(Ordering::Relaxed)
2187}
2188pub(crate) fn mm_kill() {
2189    MM_KILL.store(true, Ordering::Relaxed);
2190}
2191
2192/// Consecutive over-budget ops. ONE slow op is not contention: on a
2193/// 24 GB Mac running the 25.7 GB fl2va file the first ops after the
2194/// prompt encode page their weights in from the SSD and take seconds —
2195/// a field report (hololabs, HF discussion #2) had to neuter the kill
2196/// to keep the denoise on the GPU, and then measured 48 s/step where the
2197/// CPU fallback took >60. Contention is persistent; a page-in is not.
2198static MM_STRIKES: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
2199const MM_STRIKES_TO_KILL: u32 = 3;
2200/// Whether the kill is armed at all. A one-shot phase whose slowness is
2201/// expected and not contention — the video prompt encoder streaming
2202/// 12 GB off the SSD on a 24 GB Mac (HF discussion #4: users had to
2203/// gut `mm_kill` to keep the denoise loop on the GPU) — disarms it and
2204/// re-arms it when the phase is over; strikes taken meanwhile are
2205/// forgotten.
2206static MM_ARMED: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
2207
2208/// Disarm / re-arm the contention kill around a phase whose GEMMs are
2209/// slow for reasons that are not another process (see `MM_ARMED`).
2210pub fn mm_kill_arm(on: bool) {
2211    MM_ARMED.store(on, Ordering::Relaxed);
2212    if on {
2213        MM_STRIKES.store(0, Ordering::Relaxed);
2214    }
2215}
2216
2217/// The contention verdict for one wide op: `el` against its
2218/// work-proportional `budget`. `exempt` marks ops whose time is not
2219/// evidence — the cold probe, or a weight that was not resident before
2220/// the call and rode in with it. Kills after `MM_STRIKES_TO_KILL`
2221/// consecutive strikes; a within-budget op clears the count.
2222/// `CMF_MM_KILL=0` disables the kill entirely (the device is trusted).
2223pub(crate) fn mm_budget_check(
2224    what: &str,
2225    el: std::time::Duration,
2226    budget: std::time::Duration,
2227    exempt: bool,
2228) {
2229    if el <= budget {
2230        MM_STRIKES.store(0, Ordering::Relaxed);
2231        return;
2232    }
2233    if exempt || !MM_ARMED.load(Ordering::Relaxed) {
2234        return;
2235    }
2236    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2237    let on = *ON.get_or_init(|| std::env::var("CMF_MM_KILL").as_deref() != Ok("0"));
2238    let n = MM_STRIKES.fetch_add(1, Ordering::Relaxed) + 1;
2239    if !on {
2240        tracing::info!(
2241            "gpu {what} took {el:?} (budget {budget:?}) — over budget, CMF_MM_KILL=0 keeps the device"
2242        );
2243        return;
2244    }
2245    if n >= MM_STRIKES_TO_KILL {
2246        tracing::warn!(
2247            "gpu {what} took {el:?} (budget {budget:?}), {n} in a row — \
2248             device contended, CPU for the rest of the process (CMF_MM_KILL=0 to override)"
2249        );
2250        mm_kill();
2251    } else {
2252        tracing::info!(
2253            "gpu {what} took {el:?} (budget {budget:?}) — strike {n} of {MM_STRIKES_TO_KILL}"
2254        );
2255    }
2256}
2257
2258/// Fused DiT SwiGLU FFN on the device: g=X·W1ᵀ, u=X·W3ᵀ, silu(g)·u,
2259/// Causal chunk attention on the device: `b` queries against `s0 + b`
2260/// cached keys. wgpu only — Metal's chunk graph keeps attention inside
2261/// the resident block and never calls out.
2262#[allow(unused_variables, clippy::too_many_arguments)]
2263pub fn chunk_attend(
2264    q: &[f32],
2265    k: &[&[f32]],
2266    v: &[&[f32]],
2267    b: usize,
2268    s0: usize,
2269    nh: usize,
2270    nkv: usize,
2271    hd: usize,
2272    scale: f32,
2273    out: &mut [f32],
2274) -> bool {
2275    match backend() {
2276        #[cfg(feature = "gpu")]
2277        Backend::Wgpu => crate::gpu_wgpu::chunk_attend(q, k, v, b, s0, nh, nkv, hd, scale, out),
2278        #[allow(unreachable_patterns)]
2279        _ => false,
2280    }
2281}
2282
2283/// Fused QKV projection: one upload of the normed chunk, three GEMMs,
2284/// one readback of Q|K|V back to back. Metal has no twin yet — its
2285/// chunk graph keeps the whole layer resident and never surfaces QKV.
2286#[allow(unused_variables, clippy::too_many_arguments)]
2287pub fn q4t_qkv(
2288    model: &Arc<CmfModel>,
2289    wq: usize,
2290    wk: usize,
2291    wv: usize,
2292    xs: &[f32],
2293    b: usize,
2294    cols: usize,
2295    rq: usize,
2296    rk: usize,
2297    rv: usize,
2298    out: &mut [f32],
2299) -> bool {
2300    match backend() {
2301        #[cfg(feature = "gpu")]
2302        Backend::Wgpu => crate::gpu_wgpu::q4t_qkv(model, wq, wk, wv, xs, b, cols, rq, rk, rv, out),
2303        #[allow(unreachable_patterns)]
2304        _ => false,
2305    }
2306}
2307
2308/// y=·W2ᵀ — one command buffer, only X and Y cross the CPU boundary.
2309#[allow(unused_variables, clippy::too_many_arguments)]
2310/// SwiGLU FFN with a row-packed [gate|up] fc1 (MiniMax-H3's DiT), run
2311/// end to end on the device. wgpu only: Metal keeps the host loop until
2312/// its own packed kernel exists.
2313#[allow(clippy::too_many_arguments, unused_variables)]
2314pub fn q4tp_ffn_packed(
2315    model: &Arc<CmfModel>,
2316    w1: usize,
2317    w2: usize,
2318    xs: &[f32],
2319    b: usize,
2320    hidden: usize,
2321    inter: usize,
2322    bias: Option<&[f32]>,
2323    out: &mut [f32],
2324) -> bool {
2325    match backend() {
2326        #[cfg(feature = "gpu")]
2327        Backend::Wgpu => {
2328            crate::gpu_wgpu::ffn_packed(model, w1, w2, xs, b, hidden, inter, bias, out)
2329        }
2330        #[allow(unreachable_patterns)]
2331        _ => false,
2332    }
2333}
2334
2335pub fn q4tp_ffn(
2336    model: &Arc<CmfModel>,
2337    w1: usize,
2338    w3: usize,
2339    w2: usize,
2340    xs: &[f32],
2341    b: usize,
2342    hidden: usize,
2343    inter: usize,
2344    out: &mut [f32],
2345) -> bool {
2346    match backend() {
2347        #[cfg(target_os = "macos")]
2348        Backend::Metal => crate::gpu_metal::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2349        #[cfg(feature = "gpu")]
2350        Backend::Wgpu => crate::gpu_wgpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2351        #[allow(unreachable_patterns)]
2352        _ => false,
2353    }
2354}
2355
2356/// Qwen Image's exact two-projection tanh-GELU FFN.  The WGPU arm keeps the
2357/// intermediate on the device; other backends decline so the caller retains
2358/// its bounded CPU path.  `bias_in` is applied before GELU and `bias_out`
2359/// after the second projection, matching the official transformer.
2360#[allow(clippy::too_many_arguments, unused_variables)]
2361pub fn q4tp_gelu_ffn(
2362    model: &Arc<CmfModel>,
2363    w_in: usize,
2364    w_out: usize,
2365    xs: &[f32],
2366    b: usize,
2367    hidden: usize,
2368    inter: usize,
2369    bias_in: &[f32],
2370    bias_out: &[f32],
2371    out: &mut [f32],
2372) -> bool {
2373    match backend() {
2374        #[cfg(feature = "gpu")]
2375        Backend::Wgpu => crate::gpu_wgpu::q4tp_gelu_ffn(
2376            model, w_in, w_out, xs, b, hidden, inter, bias_in, bias_out, out,
2377        ),
2378        #[allow(unreachable_patterns)]
2379        _ => false,
2380    }
2381}
2382
2383pub fn q4t_ffn(
2384    model: &Arc<CmfModel>,
2385    w1: usize,
2386    w3: usize,
2387    w2: usize,
2388    xs: &[f32],
2389    b: usize,
2390    hidden: usize,
2391    inter: usize,
2392    out: &mut [f32],
2393) -> bool {
2394    match backend() {
2395        #[cfg(target_os = "macos")]
2396        Backend::Metal => crate::gpu_metal::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2397        #[cfg(feature = "gpu")]
2398        Backend::Wgpu => crate::gpu_wgpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2399        #[allow(unreachable_patterns)]
2400        _ => false,
2401    }
2402}
2403
2404/// One whole modulated DiT block for `dit_block`: geometry, norm
2405/// weights, AdaLN scale/gate vectors (gates pre-tanh'd), a per-token
2406/// f32 RoPE cos/sin table, and the directory indices of the seven
2407/// q4t projections. `x` is in-out `[n, hidden]`.
2408pub struct DitBlockArgs<'a> {
2409    pub n: usize,
2410    pub hidden: usize,
2411    pub inter: usize,
2412    pub nh: usize,
2413    pub nkv: usize,
2414    pub hd: usize,
2415    pub eps: f32,
2416    pub rope_cos: &'a [f32],
2417    pub rope_sin: &'a [f32],
2418    pub norm1: &'a [f32],
2419    pub norm2: &'a [f32],
2420    pub ffn_norm1: &'a [f32],
2421    pub ffn_norm2: &'a [f32],
2422    pub norm_q: &'a [f32],
2423    pub norm_k: &'a [f32],
2424    pub s_msa: &'a [f32],
2425    pub gate_msa: &'a [f32],
2426    pub s_mlp: &'a [f32],
2427    pub gate_mlp: &'a [f32],
2428    pub wq: usize,
2429    pub wk: usize,
2430    pub wv: usize,
2431    pub wo: usize,
2432    pub w1: usize,
2433    pub w3: usize,
2434    pub w2: usize,
2435    /// The projections' layout: q4tp (ladder scales) vs plain q4_tiled.
2436    /// The recommended Lumina file is q4tp, and a backend that only
2437    /// knows q4t must decline rather than decode with the wrong reader.
2438    pub q4tp: bool,
2439    /// The hidden state is already on the device from the previous block,
2440    /// so `x` need not be uploaded.
2441    pub resident_in: bool,
2442    /// Leave the result on the device instead of reading it back. The DiT
2443    /// loop does not touch `x` between blocks, so 27 of every 28 readbacks
2444    /// were moving 19 MB across PCIe and stalling on it for nothing.
2445    pub resident_out: bool,
2446}
2447
2448/// Can the selected backend keep the DiT's hidden state on the device
2449/// between blocks? Only the wgpu whole-block path; the Metal entry takes
2450/// and returns host memory every call.
2451pub fn dit_chain_supported() -> bool {
2452    #[cfg(feature = "gpu")]
2453    {
2454        return matches!(backend(), Backend::Wgpu) && fused_dit_block_available();
2455    }
2456    #[allow(unreachable_code)]
2457    false
2458}
2459
2460/// Pull the resident hidden state back to the host. For the caller that
2461/// chained blocks and then hit one the device declined.
2462pub fn dit_state_fetch(_x: &mut [f32]) -> bool {
2463    #[cfg(feature = "gpu")]
2464    {
2465        if matches!(backend(), Backend::Wgpu) {
2466            return crate::gpu_wgpu::dit_state_fetch(_x);
2467        }
2468    }
2469    false
2470}
2471
2472/// One whole modulated DiT block on the device — norms, qkv, RoPE,
2473/// attention, residuals and the SwiGLU FFN in a single command
2474/// buffer; only `x` crosses the CPU boundary (in and out).
2475#[allow(unused_variables)]
2476/// The DiT's three projections in one submission (wgpu only; the
2477/// Metal path fuses the whole block instead). False = the caller keeps
2478/// its three separate calls.
2479#[allow(unused_variables, clippy::too_many_arguments)]
2480pub fn dit_qkv(
2481    model: &Arc<CmfModel>,
2482    wq: usize,
2483    wk: usize,
2484    wv: usize,
2485    xs: &[f32],
2486    b: usize,
2487    hidden: usize,
2488    qrows: usize,
2489    kvrows: usize,
2490    q_out: &mut [f32],
2491    k_out: &mut [f32],
2492    v_out: &mut [f32],
2493) -> bool {
2494    match backend() {
2495        #[cfg(feature = "gpu")]
2496        Backend::Wgpu => crate::gpu_wgpu::q4tp_qkv(
2497            model, wq, wk, wv, xs, b, hidden, qrows, kvrows, q_out, k_out, v_out,
2498        ),
2499        #[allow(unreachable_patterns)]
2500        _ => false,
2501    }
2502}
2503
2504/// The Qwen Image double-stream attention half.  The WGPU implementation
2505/// keeps the six Q/K/V projections, the stream join, qk-norm/RoPE, joint
2506/// attention, and both output projections on the device; the caller only
2507/// supplies the two normalized streams and receives the two projected
2508/// streams.  A backend or codec that cannot satisfy the full contract
2509/// returns `false` before changing either output, so the native host path
2510/// remains the portable fallback.
2511pub struct QwenImageAttentionArgs<'a> {
2512    pub image: &'a [f32],
2513    pub text: &'a [f32],
2514    pub image_tokens: usize,
2515    pub text_tokens: usize,
2516    pub heads: usize,
2517    pub head_dim: usize,
2518    pub image_q: usize,
2519    pub image_k: usize,
2520    pub image_v: usize,
2521    pub text_q: usize,
2522    pub text_k: usize,
2523    pub text_v: usize,
2524    pub image_out: usize,
2525    pub text_out: usize,
2526    pub image_q_norm: &'a [f32],
2527    pub image_k_norm: &'a [f32],
2528    pub text_q_norm: &'a [f32],
2529    pub text_k_norm: &'a [f32],
2530    pub image_cos: &'a [f32],
2531    pub image_sin: &'a [f32],
2532    pub text_cos: &'a [f32],
2533    pub text_sin: &'a [f32],
2534    pub image_q_bias: &'a [f32],
2535    pub image_k_bias: &'a [f32],
2536    pub image_v_bias: &'a [f32],
2537    pub text_q_bias: &'a [f32],
2538    pub text_k_bias: &'a [f32],
2539    pub text_v_bias: &'a [f32],
2540    pub image_out_bias: &'a [f32],
2541    pub text_out_bias: &'a [f32],
2542    pub image_proj: &'a mut [f32],
2543    pub text_proj: &'a mut [f32],
2544}
2545
2546/// The per-layer controls and Q4TP directory indices used by the native
2547/// Qwen block.  Keeping this descriptor separate from the stream buffers
2548/// lets a whole transformer forward reuse one explicit device state without
2549/// a global scratch slot or a hidden context label.
2550#[allow(clippy::too_many_fields)]
2551pub struct QwenImageChainBlock<'a> {
2552    pub image_mod: &'a [f32],
2553    pub text_mod: &'a [f32],
2554    pub image_q: usize,
2555    pub image_k: usize,
2556    pub image_v: usize,
2557    pub text_q: usize,
2558    pub text_k: usize,
2559    pub text_v: usize,
2560    pub image_out: usize,
2561    pub text_out: usize,
2562    pub image_q_norm: &'a [f32],
2563    pub image_k_norm: &'a [f32],
2564    pub text_q_norm: &'a [f32],
2565    pub text_k_norm: &'a [f32],
2566    pub image_q_bias: &'a [f32],
2567    pub image_k_bias: &'a [f32],
2568    pub image_v_bias: &'a [f32],
2569    pub text_q_bias: &'a [f32],
2570    pub text_k_bias: &'a [f32],
2571    pub text_v_bias: &'a [f32],
2572    pub image_out_bias: &'a [f32],
2573    pub text_out_bias: &'a [f32],
2574    pub image_attn_gate: &'a [f32],
2575    pub text_attn_gate: &'a [f32],
2576    pub image_mlp_in: usize,
2577    pub image_mlp_out: usize,
2578    pub text_mlp_in: usize,
2579    pub text_mlp_out: usize,
2580    pub image_mlp_in_bias: &'a [f32],
2581    pub image_mlp_out_bias: &'a [f32],
2582    pub text_mlp_in_bias: &'a [f32],
2583    pub text_mlp_out_bias: &'a [f32],
2584}
2585
2586/// Complete Qwen Image transformer block contract. The first norm/mod
2587/// panels are supplied by the native caller; the WGPU arm keeps both streams
2588/// resident through QKV, QK/RoPE, joint attention, output projections, both
2589/// gated residuals, and the exact tanh-GELU MLPs. A backend that cannot
2590/// satisfy the whole graph returns `false` without changing either output.
2591#[allow(clippy::too_many_fields)]
2592pub struct QwenImageBlockArgs<'a> {
2593    /// Raw stream state is read for the first gated residual and overwritten
2594    /// with the block's final state after the one readback.
2595    pub image: &'a mut [f32],
2596    pub text: &'a mut [f32],
2597    pub image_norm: &'a [f32],
2598    pub text_norm: &'a [f32],
2599    pub image_tokens: usize,
2600    pub text_tokens: usize,
2601    pub heads: usize,
2602    pub head_dim: usize,
2603    pub image_cos: &'a [f32],
2604    pub image_sin: &'a [f32],
2605    pub text_cos: &'a [f32],
2606    pub text_sin: &'a [f32],
2607    pub image_q: usize,
2608    pub image_k: usize,
2609    pub image_v: usize,
2610    pub text_q: usize,
2611    pub text_k: usize,
2612    pub text_v: usize,
2613    pub image_out: usize,
2614    pub text_out: usize,
2615    pub image_q_norm: &'a [f32],
2616    pub image_k_norm: &'a [f32],
2617    pub text_q_norm: &'a [f32],
2618    pub text_k_norm: &'a [f32],
2619    pub image_q_bias: &'a [f32],
2620    pub image_k_bias: &'a [f32],
2621    pub image_v_bias: &'a [f32],
2622    pub text_q_bias: &'a [f32],
2623    pub text_k_bias: &'a [f32],
2624    pub text_v_bias: &'a [f32],
2625    pub image_out_bias: &'a [f32],
2626    pub text_out_bias: &'a [f32],
2627    pub image_attn_gate: &'a [f32],
2628    pub text_attn_gate: &'a [f32],
2629    pub image_mlp_in: usize,
2630    pub image_mlp_out: usize,
2631    pub text_mlp_in: usize,
2632    pub text_mlp_out: usize,
2633    pub image_mlp_in_bias: &'a [f32],
2634    pub image_mlp_out_bias: &'a [f32],
2635    pub text_mlp_in_bias: &'a [f32],
2636    pub text_mlp_out_bias: &'a [f32],
2637    pub image_mlp_mod: &'a [f32],
2638    pub text_mlp_mod: &'a [f32],
2639    pub image_mlp_gate: &'a [f32],
2640    pub text_mlp_gate: &'a [f32],
2641}
2642
2643/// Explicit whole-forward Qwen state contract.  The WGPU backend uploads the
2644/// two initial streams once, encodes a bounded number of complete blocks per
2645/// submission, and reads the final state once.  `blocks` is immutable for the
2646/// call, while the two stream slices receive only the final readback.
2647pub struct QwenImageChainArgs<'a> {
2648    pub image: &'a mut [f32],
2649    pub text: &'a mut [f32],
2650    pub image_tokens: usize,
2651    pub text_tokens: usize,
2652    pub heads: usize,
2653    pub head_dim: usize,
2654    pub image_cos: &'a [f32],
2655    pub image_sin: &'a [f32],
2656    pub text_cos: &'a [f32],
2657    pub text_sin: &'a [f32],
2658    pub blocks: &'a [QwenImageChainBlock<'a>],
2659}
2660
2661#[allow(unused_variables)]
2662pub fn qwen_image_attention(
2663    model: &Arc<CmfModel>,
2664    args: &mut QwenImageAttentionArgs<'_>,
2665) -> bool {
2666    match backend() {
2667        #[cfg(feature = "gpu")]
2668        Backend::Wgpu => crate::gpu_wgpu::qwen_image_attention(model, args),
2669        #[allow(unreachable_patterns)]
2670        _ => false,
2671    }
2672}
2673
2674#[allow(unused_variables)]
2675pub fn qwen_image_block(model: &Arc<CmfModel>, args: &mut QwenImageBlockArgs<'_>) -> bool {
2676    match backend() {
2677        #[cfg(feature = "gpu")]
2678        Backend::Wgpu => crate::gpu_wgpu::qwen_image_block(model, args),
2679        #[allow(unreachable_patterns)]
2680        _ => false,
2681    }
2682}
2683
2684/// Keep all Qwen transformer blocks on the selected WGPU device, with only
2685/// bounded chunk submissions and one final readback.  Other backends decline
2686/// so the native caller can use its exact portable block loop.
2687#[allow(unused_variables)]
2688pub fn qwen_image_chain(model: &Arc<CmfModel>, args: &mut QwenImageChainArgs<'_>) -> bool {
2689    match backend() {
2690        #[cfg(feature = "gpu")]
2691        Backend::Wgpu => crate::gpu_wgpu::qwen_image_chain(model, args),
2692        #[allow(unreachable_patterns)]
2693        _ => false,
2694    }
2695}
2696
2697/// The Qwen Image second sub-block on WGPU: affine-free LayerNorm,
2698/// shift/scale modulation, Q4TP input projection, exact tanh-GELU, output
2699/// projection, bias and gated residual.  `data` is updated in place after a
2700/// single final readback.  Backends/codecs that cannot keep this chain on the
2701/// device return `false` before changing `data`, leaving the caller's
2702/// portable per-op path intact.
2703#[allow(unused_variables, clippy::too_many_arguments)]
2704pub fn qwen_image_mlp_inplace(
2705    model: &Arc<CmfModel>,
2706    w_in: usize,
2707    w_out: usize,
2708    data: &mut [f32],
2709    batch: usize,
2710    hidden: usize,
2711    inter: usize,
2712    bias_in: &[f32],
2713    bias_out: &[f32],
2714    modulation: &[f32],
2715    gate: &[f32],
2716) -> bool {
2717    match backend() {
2718        #[cfg(feature = "gpu")]
2719        Backend::Wgpu => crate::gpu_wgpu::qwen_image_mlp_inplace(
2720            model,
2721            w_in,
2722            w_out,
2723            data,
2724            batch,
2725            hidden,
2726            inter,
2727            bias_in,
2728            bias_out,
2729            modulation,
2730            gate,
2731        ),
2732        #[allow(unreachable_patterns)]
2733        _ => false,
2734    }
2735}
2736
2737/// Is a FUSED whole-block device path on offer? The batched-CFG shape
2738/// (two sequences in one tall batch) and the fused block (one sequence,
2739/// one command buffer) are alternatives, and the caller picks.
2740pub fn fused_dit_block_available() -> bool {
2741    #[cfg(target_os = "macos")]
2742    {
2743        matches!(backend(), Backend::Metal) && fused_block_trusted()
2744    }
2745    #[cfg(not(target_os = "macos"))]
2746    {
2747        false
2748    }
2749}
2750
2751pub fn dit_block(model: &Arc<CmfModel>, a: &DitBlockArgs, x: &mut [f32]) -> bool {
2752    dit_block_seg(model, a, &[a.n], x)
2753}
2754
2755/// The same block over a CONCATENATION of independent sequences:
2756/// attention per segment, everything position-wise batched. wgpu only —
2757/// the Metal path takes the single-sequence entry above.
2758pub fn dit_block_seg(
2759    model: &Arc<CmfModel>,
2760    a: &DitBlockArgs,
2761    segs: &[usize],
2762    x: &mut [f32],
2763) -> bool {
2764    match backend() {
2765        #[cfg(target_os = "macos")]
2766        Backend::Metal if segs.len() <= 1 => crate::gpu_metal::dit_block(model, a, x),
2767        // The wgpu whole-block path. What it buys is host round trips —
2768        // six a block become one — so it defaults ON where those cost
2769        // real time (a discrete card across PCIe) and OFF on unified
2770        // memory, where the per-op path shares the same pages and the
2771        // fusion measured slightly slower on an M4. `CMF_DIT_FUSED=1`
2772        // forces it anywhere, `=0` forbids it.
2773        #[cfg(feature = "gpu")]
2774        Backend::Wgpu
2775            if match std::env::var("CMF_DIT_FUSED").ok().as_deref() {
2776                Some("0") => false,
2777                Some(_) => true,
2778                None => crate::gpu_wgpu::discrete_active(),
2779            } =>
2780        {
2781            crate::gpu_wgpu::dit_block_seg(model, a, segs, x)
2782        }
2783        #[allow(unreachable_patterns)]
2784        _ => false,
2785    }
2786}
2787
2788/// One VAE resnet block for `vae_resnet`: norm/conv weights and the
2789/// channel/shape geometry. `shortcut` is the 1×1 projection (w, b, k)
2790/// when in/out channels differ.
2791pub struct VaeResnetArgs<'a> {
2792    pub groups: usize,
2793    pub ic: usize,
2794    pub oc: usize,
2795    pub h: usize,
2796    pub w: usize,
2797    pub n1w: &'a [f32],
2798    pub n1b: &'a [f32],
2799    pub c1w: &'a [f32],
2800    pub c1b: &'a [f32],
2801    pub c1k: usize,
2802    pub n2w: &'a [f32],
2803    pub n2b: &'a [f32],
2804    pub c2w: &'a [f32],
2805    pub c2b: &'a [f32],
2806    pub c2k: usize,
2807    pub shortcut: Option<(&'a [f32], &'a [f32], usize)>,
2808}
2809
2810/// One whole VAE resnet block on the device (norm+silu → conv ×2 →
2811/// shortcut → add, one command buffer).
2812#[allow(unused_variables)]
2813pub fn vae_resnet(a: &VaeResnetArgs, x: &[f32], out: &mut [f32]) -> bool {
2814    match backend() {
2815        #[cfg(target_os = "macos")]
2816        Backend::Metal => crate::gpu_metal::vae_resnet(a, x, out),
2817        _ => false,
2818    }
2819}
2820
2821/// Nearest-2× upsample fused with the following conv — the small
2822/// pre-upsample image is what crosses the CPU boundary.
2823#[allow(unused_variables, clippy::too_many_arguments)]
2824pub fn vae_upsample_conv(
2825    w: &[f32],
2826    bias: &[f32],
2827    x: &[f32],
2828    ic: usize,
2829    oc: usize,
2830    h: usize,
2831    w_img: usize,
2832    k: usize,
2833    out: &mut [f32],
2834) -> bool {
2835    match backend() {
2836        #[cfg(target_os = "macos")]
2837        Backend::Metal => crate::gpu_metal::vae_upsample_conv(w, bias, x, ic, oc, h, w_img, k, out),
2838        #[cfg(feature = "gpu")]
2839        Backend::Wgpu => crate::gpu_wgpu::vae_upsample_conv(w, bias, x, ic, oc, h, w_img, k, out),
2840        #[allow(unreachable_patterns)]
2841        _ => false,
2842    }
2843}
2844
2845/// VAE conv2d on the device (implicit GEMM — the CPU path pays for a
2846/// multi-GB im2col matrix at high resolutions).
2847#[allow(unused_variables, clippy::too_many_arguments)]
2848pub fn vae_conv2d(
2849    w: &[f32],
2850    bias: &[f32],
2851    x: &[f32],
2852    ic: usize,
2853    oc: usize,
2854    h: usize,
2855    w_img: usize,
2856    k: usize,
2857    out: &mut [f32],
2858) -> bool {
2859    match backend() {
2860        #[cfg(target_os = "macos")]
2861        Backend::Metal => crate::gpu_metal::vae_conv2d(w, bias, x, ic, oc, h, w_img, k, out),
2862        #[cfg(feature = "gpu")]
2863        Backend::Wgpu => crate::gpu_wgpu::vae_conv2d(w, bias, x, ic, oc, h, w_img, k, out),
2864        #[allow(unreachable_patterns)]
2865        _ => false,
2866    }
2867}
2868
2869/// DiT full bidirectional attention on the device (all heads:
2870/// scores GEMM → row softmax → P·V → panel unstack, one command
2871/// buffer). Head-major inputs; out is [n, nh·hd].
2872#[allow(unused_variables, clippy::too_many_arguments)]
2873/// Attention from an interleaved qkv panel, splitting into head-major
2874/// planes ON the device. wgpu only; `false` elsewhere so the caller
2875/// keeps its host repack.
2876#[allow(unused_variables)]
2877#[allow(clippy::too_many_arguments)]
2878/// qkv projection + attention with the panel never leaving the card.
2879/// wgpu only; `false` elsewhere and the caller keeps its host chain.
2880#[allow(clippy::too_many_arguments, unused_variables)]
2881pub fn dit_qkv_attention(
2882    model: &Arc<CmfModel>,
2883    qkv_idx: usize,
2884    xn: &[f32],
2885    n: usize,
2886    hidden: usize,
2887    nh: usize,
2888    hd: usize,
2889    scale: f32,
2890    nr: (&[f32], &[f32], &[f32], f32),
2891    out: &mut [f32],
2892) -> bool {
2893    match backend() {
2894        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2895        Backend::Wgpu => crate::gpu_wgpu::dit_qkv_attention(
2896            model, qkv_idx, xn, n, hidden, nh, hd, scale, nr, out,
2897        ),
2898        #[allow(unreachable_patterns)]
2899        _ => false,
2900    }
2901}
2902
2903/// The whole attention half of a DiT block on the card: qkv GEMM,
2904/// attention, output projection. Only `proj` comes home.
2905#[allow(clippy::too_many_arguments)]
2906pub fn dit_qkv_attn_out(
2907    model: &Arc<CmfModel>,
2908    qkv_idx: usize,
2909    out_idx: usize,
2910    xn: &[f32],
2911    n: usize,
2912    hidden: usize,
2913    nh: usize,
2914    hd: usize,
2915    scale: f32,
2916    nr: (&[f32], &[f32], &[f32], f32),
2917    proj: &mut [f32],
2918) -> bool {
2919    match backend() {
2920        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2921        Backend::Wgpu => crate::gpu_wgpu::dit_qkv_attn_out(
2922            model, qkv_idx, out_idx, xn, n, hidden, nh, hd, scale, nr, proj,
2923        ),
2924        #[allow(unreachable_patterns)]
2925        _ => false,
2926    }
2927}
2928
2929/// The VAE decoder's attention half on the card. Only `proj` returns.
2930#[allow(clippy::too_many_arguments)]
2931pub fn vae_qkv_attn_out(
2932    model: &Arc<CmfModel>,
2933    qkv_idx: usize,
2934    out_idx: usize,
2935    xn: &[f32],
2936    n: usize,
2937    dim: usize,
2938    nh: usize,
2939    hd: usize,
2940    scale: f32,
2941    angles: &[f32],
2942    eps: f32,
2943    qkv_bias: &[f32],
2944    proj: &mut [f32],
2945) -> bool {
2946    match backend() {
2947        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2948        Backend::Wgpu => crate::gpu_wgpu::vae_qkv_attn_out(
2949            model, qkv_idx, out_idx, xn, n, dim, nh, hd, scale, angles, eps, qkv_bias, proj,
2950        ),
2951        #[allow(unreachable_patterns)]
2952        _ => false,
2953    }
2954}
2955
2956#[allow(clippy::too_many_arguments)]
2957pub fn vae_attention_packed(
2958    qkv: &[f32],
2959    nh: usize,
2960    n: usize,
2961    hd: usize,
2962    scale: f32,
2963    angles: &[f32],
2964    eps: f32,
2965    out: &mut [f32],
2966) -> bool {
2967    vae_attention_packed_layout(qkv, nh, n, hd, scale, angles, eps, out, 1)
2968}
2969
2970#[allow(clippy::too_many_arguments)]
2971pub fn vae_attention_packed_layout(
2972    qkv: &[f32],
2973    nh: usize,
2974    n: usize,
2975    hd: usize,
2976    scale: f32,
2977    angles: &[f32],
2978    eps: f32,
2979    out: &mut [f32],
2980    layout: u32,
2981) -> bool {
2982    match backend() {
2983        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
2984        Backend::Wgpu => crate::gpu_wgpu::vae_attention_packed_layout(
2985            qkv, nh, n, hd, scale, angles, eps, out, layout,
2986        ),
2987        #[allow(unreachable_patterns)]
2988        _ => false,
2989    }
2990}
2991
2992#[allow(clippy::too_many_arguments)]
2993pub fn dit_split_only(
2994    qkv: &[f32],
2995    nh: usize,
2996    n: usize,
2997    hd: usize,
2998    layout: u32,
2999    norm: Option<(&[f32], f32)>,
3000    out_q: &mut [f32],
3001) -> bool {
3002    match backend() {
3003        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3004        Backend::Wgpu => crate::gpu_wgpu::dit_split_only(qkv, nh, n, hd, layout, norm, out_q),
3005        #[allow(unreachable_patterns)]
3006        _ => false,
3007    }
3008}
3009
3010/// The backend's f32 NT GEMM: `y[n×m] = x[n×k] · wᵀ[m×k]`. Tensor
3011/// cores where the card has them. Refuses under `CMF_BAKE_GPU=0` or
3012/// strict f32, and for jobs below n·k·m = 4M, where the round trip
3013/// costs more than the arithmetic saves.
3014/// `gemm_nt_f32` whose `w` is known to change every call (an
3015/// accumulation over fresh activations, not a weight): it skips the
3016/// resident ledger and its per-call fingerprint of the whole operand.
3017pub fn gemm_nt_f32_transient(
3018    x: &[f32],
3019    w: &[f32],
3020    y: &mut [f32],
3021    n: usize,
3022    k: usize,
3023    m: usize,
3024) -> bool {
3025    match backend() {
3026        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3027        Backend::Wgpu => crate::gpu_wgpu::gemm_nt_f32_transient(x, w, y, n, k, m),
3028        #[allow(unreachable_patterns)]
3029        _ => false,
3030    }
3031}
3032
3033pub fn gemm_nt_f32(x: &[f32], w: &[f32], y: &mut [f32], n: usize, k: usize, m: usize) -> bool {
3034    match backend() {
3035        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3036        Backend::Wgpu => crate::gpu_wgpu::gemm_nt_f32(x, w, y, n, k, m),
3037        #[allow(unreachable_patterns)]
3038        _ => false,
3039    }
3040}
3041
3042/// Music-3's FFN chain resident on the device — two GEMMs and the GLU
3043/// between them with no host round trip. `false` = refused, host runs.
3044#[allow(clippy::too_many_arguments)]
3045pub fn music3_ffn(
3046    model: &std::sync::Arc<CmfModel>,
3047    idx_in: usize,
3048    idx_out: usize,
3049    h: &[f32],
3050    bias_in: &[f32],
3051    n: usize,
3052    hs: usize,
3053    inter: usize,
3054    out: &mut [f32],
3055) -> bool {
3056    match backend() {
3057        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3058        Backend::Wgpu => {
3059            crate::gpu_wgpu::music3_ffn(model, idx_in, idx_out, h, bias_in, n, hs, inter, out)
3060        }
3061        #[allow(unreachable_patterns)]
3062        _ => false,
3063    }
3064}
3065
3066/// A 1D convolution as a GEMM whose column matrix is expanded on the
3067/// device instead of being built, transposed and uploaded by the host.
3068/// `yt` comes back `[out_n x oc]`. `false` = refused, caller runs host.
3069#[allow(clippy::too_many_arguments)]
3070pub fn conv1d_gemm(
3071    x: &[f32],
3072    w: &[f32],
3073    ic: usize,
3074    oc: usize,
3075    n: usize,
3076    k: usize,
3077    pad: usize,
3078    dil: usize,
3079    out_n: usize,
3080    yt: &mut [f32],
3081) -> bool {
3082    match backend() {
3083        #[cfg(target_os = "macos")]
3084        Backend::Metal => crate::gpu_metal::conv1d_gemm(x, w, ic, oc, n, k, pad, dil, out_n, yt),
3085        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3086        Backend::Wgpu => crate::gpu_wgpu::conv1d_gemm(x, w, ic, oc, n, k, pad, dil, out_n, yt),
3087        #[allow(unreachable_patterns)]
3088        _ => false,
3089    }
3090}
3091
3092/// The convolution as a GEMM on the matrix units. `false` = refused.
3093#[allow(clippy::too_many_arguments)]
3094pub fn vae_conv2d_coop(
3095    w: &[f32],
3096    bias: Option<&[f32]>,
3097    x: &[f32],
3098    ic: usize,
3099    oc: usize,
3100    h: usize,
3101    wi: usize,
3102    k: usize,
3103    out: &mut [f32],
3104) -> bool {
3105    match backend() {
3106        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3107        Backend::Wgpu => crate::gpu_wgpu::vae_conv2d_coop(w, bias, x, ic, oc, h, wi, k, out),
3108        #[allow(unreachable_patterns)]
3109        _ => false,
3110    }
3111}
3112
3113pub fn dit_attention_packed(
3114    qkv: &[f32],
3115    nh: usize,
3116    n: usize,
3117    hd: usize,
3118    scale: f32,
3119    // (rope angles, q norm weights, k norm weights, eps) when the device
3120    // should apply qk-norm and RoPE itself; None when the host already did.
3121    nr: Option<(&[f32], &[f32], &[f32], f32)>,
3122    out: &mut [f32],
3123) -> bool {
3124    match backend() {
3125        // wgpu carries the only implementation, and it is not
3126        // platform-specific: `CMF_GPU=wgpu` on macOS runs it over Metal
3127        // like anywhere else. It used to be compiled out here on macOS,
3128        // which made the call a silent `false` — and the caller's
3129        // `assert!` turned that refusal into a panic on every
3130        // `cortiq animate` this platform ever ran.
3131        #[cfg(feature = "gpu")]
3132        Backend::Wgpu => crate::gpu_wgpu::dit_attention_packed(qkv, nh, n, hd, scale, nr, out),
3133        #[allow(unreachable_patterns)]
3134        _ => false,
3135    }
3136}
3137
3138/// Whether `dit_attention_packed` has an implementation on the backend
3139/// that is actually selected.
3140///
3141/// The caller has to know BEFORE it skips the host qk-norm: deferring
3142/// the norm to a device that then refuses leaves q/k unnormalized with
3143/// no way back. Native Metal has no packed kernel, so on macOS this is
3144/// false unless `CMF_GPU=wgpu` picked the other backend.
3145pub fn dit_attention_packed_available() -> bool {
3146    #[allow(unreachable_patterns)]
3147    match backend() {
3148        #[cfg(feature = "gpu")]
3149        Backend::Wgpu => crate::gpu_wgpu::dit_attention_packed_ready(),
3150        _ => false,
3151    }
3152}
3153
3154pub fn dit_attention(
3155    qh: &[f32],
3156    kh: &[f32],
3157    vh: &[f32],
3158    nh: usize,
3159    nkv: usize,
3160    n: usize,
3161    hd: usize,
3162    scale: f32,
3163    out: &mut [f32],
3164) -> bool {
3165    match backend() {
3166        #[cfg(target_os = "macos")]
3167        Backend::Metal => crate::gpu_metal::dit_attention(qh, kh, vh, nh, nkv, n, hd, scale, out),
3168        #[cfg(feature = "gpu")]
3169        Backend::Wgpu => crate::gpu_wgpu::dit_attention(qh, kh, vh, nh, nkv, n, hd, scale, out),
3170        #[allow(unreachable_patterns)]
3171        _ => false,
3172    }
3173}
3174
3175/// Batched q4t GEMM on the device (imagegen DiT prefill shapes).
3176/// Metal: q4t_mul_mm decodes the mmap-resident tiles inside the
3177/// GEMM's K loop. wgpu (Vulkan/DX12 → NVIDIA/AMD/Intel/Adreno/Mali):
3178/// the register-blocked WGSL twin, weights cached in VRAM.
3179#[allow(unused_variables)]
3180pub fn q4tp_matmat(
3181    model: &Arc<CmfModel>,
3182    idx: usize,
3183    xs: &[f32],
3184    b: usize,
3185    rows: usize,
3186    cols: usize,
3187    out: &mut [f32],
3188) -> bool {
3189    match backend() {
3190        #[cfg(target_os = "macos")]
3191        Backend::Metal => crate::gpu_metal::q4tp_matmat(model, idx, xs, b, rows, cols, out),
3192        #[cfg(feature = "gpu")]
3193        Backend::Wgpu => crate::gpu_wgpu::q4tp_matmat(model, idx, xs, b, rows, cols, out),
3194        #[allow(unreachable_patterns)]
3195        _ => false,
3196    }
3197}
3198
3199/// The same over a two-bit weight plane. Native Metal uses the dedicated
3200/// q2tp tile; unsupported shapes return false and preserve the host fallback.
3201pub fn q2tp_matmat(
3202    model: &Arc<CmfModel>,
3203    idx: usize,
3204    xs: &[f32],
3205    b: usize,
3206    rows: usize,
3207    cols: usize,
3208    out: &mut [f32],
3209) -> bool {
3210    match backend() {
3211        #[cfg(target_os = "macos")]
3212        Backend::Metal => crate::gpu_metal::q2tp_matmat(model, idx, xs, b, rows, cols, out),
3213        #[cfg(feature = "gpu")]
3214        Backend::Wgpu => crate::gpu_wgpu::q2tp_matmat(model, idx, xs, b, rows, cols, out),
3215        #[allow(unreachable_patterns)]
3216        _ => false,
3217    }
3218}
3219
3220/// Descriptor-aware q2tp GEMM. The affine center is selected only for a
3221/// validated q2tp_affine target; the raw dtype16 payload remains unchanged.
3222pub fn q2tp_affine_matmat(
3223    model: &Arc<CmfModel>,
3224    idx: usize,
3225    xs: &[f32],
3226    b: usize,
3227    rows: usize,
3228    cols: usize,
3229    out: &mut [f32],
3230) -> bool {
3231    match backend() {
3232        #[cfg(target_os = "macos")]
3233        Backend::Metal => crate::gpu_metal::q2tp_affine_matmat(model, idx, xs, b, rows, cols, out),
3234        #[cfg(feature = "gpu")]
3235        Backend::Wgpu => crate::gpu_wgpu::q2tp_affine_matmat(model, idx, xs, b, rows, cols, out),
3236        #[allow(unreachable_patterns)]
3237        _ => false,
3238    }
3239}
3240
3241/// Single-token q2tp matvec through the ordinary (center=1.5) WGSL kernel.
3242pub fn q2tp_matvec(
3243    model: &Arc<CmfModel>,
3244    idx: usize,
3245    xs: &[f32],
3246    rows: usize,
3247    cols: usize,
3248    out: &mut [f32],
3249) -> bool {
3250    match backend() {
3251        #[cfg(target_os = "macos")]
3252        Backend::Metal => crate::gpu_metal::q2tp_matvec(model, idx, xs, rows, cols, out),
3253        #[cfg(feature = "gpu")]
3254        Backend::Wgpu => crate::gpu_wgpu::q2tp_matvec(model, idx, xs, rows, cols, out),
3255        #[allow(unreachable_patterns)]
3256        _ => false,
3257    }
3258}
3259
3260/// Single-token q2tp matvec with the explicit affine center=1 descriptor
3261/// operator. This is kept separate from ordinary q2tp to make accidental
3262/// center changes impossible at a call site.
3263pub fn q2tp_affine_matvec(
3264    model: &Arc<CmfModel>,
3265    idx: usize,
3266    xs: &[f32],
3267    rows: usize,
3268    cols: usize,
3269    out: &mut [f32],
3270) -> bool {
3271    match backend() {
3272        #[cfg(target_os = "macos")]
3273        Backend::Metal => crate::gpu_metal::q2tp_affine_matvec(model, idx, xs, rows, cols, out),
3274        #[cfg(feature = "gpu")]
3275        Backend::Wgpu => crate::gpu_wgpu::q2tp_affine_matvec(model, idx, xs, rows, cols, out),
3276        #[allow(unreachable_patterns)]
3277        _ => false,
3278    }
3279}
3280
3281/// Single-token q4tp matvec on the device — the lm_head class. Through the
3282/// DEDICATED matvec kernel: the batched GEMM at b=1 measured 11.73 ms
3283/// against the host's 9.51 on the release head, so the route that was
3284/// supposed to save eleven milliseconds a token lost its own probe instead.
3285pub fn q4tp_matvec(
3286    model: &Arc<CmfModel>,
3287    idx: usize,
3288    xs: &[f32],
3289    rows: usize,
3290    cols: usize,
3291    out: &mut [f32],
3292) -> bool {
3293    match backend() {
3294        #[cfg(target_os = "macos")]
3295        Backend::Metal => crate::gpu_metal::q4tp_matvec_for_test(model, idx, xs, rows, cols, out),
3296        #[cfg(feature = "gpu")]
3297        Backend::Wgpu => crate::gpu_wgpu::q4tp_matvec(model, idx, xs, rows, cols, out),
3298        #[allow(unreachable_patterns)]
3299        _ => false,
3300    }
3301}
3302
3303/// Single-token q4_tiled matvec on the device — the lm_head class (a
3304/// q4t checkpoint's head is its biggest host matvec, exactly like the
3305/// q4tp twin above). wgpu holds q4t_mv pipelines only inside the graph
3306/// encoder — the standalone arm stays an honest refusal until a
3307/// discrete-GPU q4t model reaches the bench.
3308pub fn q4t_matvec(
3309    model: &Arc<CmfModel>,
3310    idx: usize,
3311    xs: &[f32],
3312    rows: usize,
3313    cols: usize,
3314    out: &mut [f32],
3315) -> bool {
3316    match backend() {
3317        #[cfg(target_os = "macos")]
3318        Backend::Metal => crate::gpu_metal::q4t_matvec_for_test(model, idx, xs, rows, cols, out),
3319        #[allow(unreachable_patterns)]
3320        _ => false,
3321    }
3322}
3323
3324pub fn q4t_matmat(
3325    model: &Arc<CmfModel>,
3326    idx: usize,
3327    xs: &[f32],
3328    b: usize,
3329    rows: usize,
3330    cols: usize,
3331    out: &mut [f32],
3332) -> bool {
3333    match backend() {
3334        #[cfg(target_os = "macos")]
3335        Backend::Metal => crate::gpu_metal::q4t_matmat(model, idx, xs, b, rows, cols, out),
3336        #[cfg(feature = "gpu")]
3337        Backend::Wgpu => crate::gpu_wgpu::q4t_matmat(model, idx, xs, b, rows, cols, out),
3338        #[allow(unreachable_patterns)]
3339        _ => false,
3340    }
3341}
3342
3343/// Whole-block token-graph types re-exported from the Metal backend.
3344#[cfg(target_os = "macos")]
3345pub use crate::gpu_metal::{
3346    AttnDeviceParams, AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GpuMoe, GraphDims, MetalFfn,
3347    O1AttnParams, TokenGraph, kv_mirror_drop, kv_mirror_read_last, kv_mirror_take_imp,
3348};
3349
3350/// A BLOCK of consecutive q1 GDN layers in one submission (Metal only).
3351#[cfg(target_os = "macos")]
3352pub fn gdn_block(
3353    model: &Arc<CmfModel>,
3354    layers: &[GdnGpuLayer],
3355    states: &mut [&mut [f32]],
3356    cfg: &GdnGpuCfg,
3357    h: &mut [f32],
3358) -> bool {
3359    match backend() {
3360        Backend::Metal => crate::gpu_metal::gdn_block(model, layers, states, cfg, h),
3361        _ => false,
3362    }
3363}
3364
3365/// A layer's MoE-FFN in one submission (amortizing the dispatch cost).
3366#[allow(unused_variables)]
3367pub fn moe_block(model: &Arc<CmfModel>, jobs: &[MoeJob], out: &mut [f32]) -> bool {
3368    match backend() {
3369        #[cfg(target_os = "macos")]
3370        Backend::Metal => crate::gpu_metal::moe_block(model, jobs, out),
3371        #[cfg(feature = "gpu")]
3372        Backend::Wgpu => crate::gpu_wgpu::moe_block(model, jobs, out),
3373        Backend::None => false,
3374    }
3375}
3376
3377/// Independent matvecs of one input in a single submission (GDN projections).
3378#[allow(unused_variables)]
3379pub fn matvec_batch(model: &Arc<CmfModel>, jobs: &[BatchJob], out: &mut [&mut [f32]]) -> bool {
3380    match backend() {
3381        #[cfg(target_os = "macos")]
3382        Backend::Metal => crate::gpu_metal::matvec_batch(model, jobs, out),
3383        #[cfg(feature = "gpu")]
3384        Backend::Wgpu => crate::gpu_wgpu::matvec_batch(model, jobs, out),
3385        Backend::None => false,
3386    }
3387}
3388
3389// ── Whole-token wgpu graph race (generation granularity) ─────────────
3390// On integrated/mobile adapters the graph is neither trusted nor banned
3391// a priori — it RACES the normal path: generations alternate arms (the
3392// normal path first — known-good UX — then the graph), per-token wall
3393// times accumulate per arm, and once both arms have enough steady
3394// samples the faster one wins for the process. Arm switches happen ONLY
3395// at generation boundaries (`kv_cache.clear()` resets state), so the
3396// device KV mirror and the CPU cache never diverge mid-sequence. The
3397// single exception is the first-token bail: the very first decode token
3398// of a graph generation may be discarded and recomputed on the CPU
3399// path (the prompt KV is CPU-owned at that point, so this is safe) —
3400// a tiled mobile GPU that drains its pipeline at every barrier turns
3401// the ~300-dispatch graph into seconds per token (field report: 0.2
3402// tok/s vs 15 on the CPU), and one token is all it takes to see that.
3403static GRAPH_RACE_STATE: AtomicU8 = AtomicU8::new(0); // 0 racing, 1 graph won, 2 normal won
3404static GRAPH_RACE_FLIP: AtomicU32 = AtomicU32::new(0);
3405static GRAPH_RACE_ARM_GRAPH: AtomicU8 = AtomicU8::new(0); // this generation's arm
3406static GRAPH_RACE_TOK: AtomicU32 = AtomicU32::new(0); // token index within the generation
3407static GRAPH_NS: [AtomicU64; 2] = [AtomicU64::new(0), AtomicU64::new(0)]; // [normal, graph]
3408static GRAPH_N: [AtomicU32; 2] = [AtomicU32::new(0), AtomicU32::new(0)];
3409
3410/// Steady per-token samples per arm before the race decides.
3411const GRAPH_RACE_SAMPLES: u32 = 4;
3412
3413// A graph that cannot be built for THIS model will never build: the
3414// refusal is a property of the weights, not of the moment. Retrying it
3415// per token is not free — the builder walks every layer and asks each
3416// tensor for a graph view before giving up at layer 0 — and on an
3417// Adreno 642L that retry cost 3x: forcing the graph on a model it
3418// refuses measured 0.3 tok/s against 0.905 for the per-op path it falls
3419// back to. Remembered once, the fallback runs at its own speed.
3420//
3421// The verdict is kept PER PIPELINE (`Pipeline::graph_refused` /
3422// `mark_graph_refused`), not in a process-wide flag: several pipelines
3423// of one file share a process in `serve` (backbone slots + skill lanes
3424// loaded mid-traffic), and a refusal in one lane — or the reset a new
3425// lane used to issue — must not move another lane's running sequence
3426// between the device and the host (R4/NF-2). Callers must NOT report
3427// transient refusals (an unsealed o1 state during prefill, a softcap).
3428
3429/// Called at every generation start (fresh KV). Applies a pending
3430/// verdict and picks this generation's arm while racing.
3431pub fn graph_race_begin_generation() {
3432    // One generation has now compiled whatever this model needs; keep it
3433    // for the next process. Once per run: the blob does not grow after
3434    // the pipelines exist, and the write is megabytes against the ~200 s
3435    // of compiling it saves on the device that needed this.
3436    #[cfg(feature = "gpu")]
3437    {
3438        // Save once, at the start of the SECOND generation: the first
3439        // has dispatched, so there is something to keep, and nothing is
3440        // saved before any work (the driver compiles at first use, not
3441        // at pipeline creation — the context comes up in 1.5 s while the
3442        // compiling costs minutes).
3443        //
3444        // Flushing again on 4, 8, 16 … was tried on the theory that a
3445        // chat turn compiles shapes the first one did not. It buys
3446        // nothing: a fresh app process still spent 49.0 s, then 58.7,
3447        // then 61.3 on its first answer with the backoff in place. One
3448        // flush it is.
3449        static FLUSHED: std::sync::Once = std::sync::Once::new();
3450        static FIRST: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
3451        if FIRST.swap(false, Ordering::Relaxed) {
3452            // Nothing dispatched yet.
3453        } else {
3454            FLUSHED.call_once(crate::gpu_wgpu::pipeline_cache_flush);
3455        }
3456    }
3457    GRAPH_RACE_TOK.store(0, Ordering::Relaxed);
3458    if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
3459        return;
3460    }
3461    let (gn, cn) = (
3462        GRAPH_N[1].load(Ordering::Relaxed),
3463        GRAPH_N[0].load(Ordering::Relaxed),
3464    );
3465    if gn >= GRAPH_RACE_SAMPLES && cn >= GRAPH_RACE_SAMPLES {
3466        let g_avg = GRAPH_NS[1].load(Ordering::Relaxed) / gn as u64;
3467        let c_avg = GRAPH_NS[0].load(Ordering::Relaxed) / cn as u64;
3468        let verdict = if g_avg < c_avg { 1 } else { 2 };
3469        GRAPH_RACE_STATE.store(verdict, Ordering::Relaxed);
3470        tracing::info!(
3471            "wgpu graph race: graph {:.2} ms/tok vs normal {:.2} ms/tok -> {}",
3472            g_avg as f64 / 1e6,
3473            c_avg as f64 / 1e6,
3474            if verdict == 1 { "graph" } else { "normal path" }
3475        );
3476        return;
3477    }
3478    let flip = GRAPH_RACE_FLIP.fetch_add(1, Ordering::Relaxed);
3479    GRAPH_RACE_ARM_GRAPH.store((flip % 2 == 1) as u8, Ordering::Relaxed);
3480}
3481
3482/// Should this decode token try the graph? `trusted` (discrete adapter,
3483/// explicit env, or a GDN hybrid whose state lives on the device) skips
3484/// the race entirely.
3485pub fn graph_race_use_graph(trusted: bool) -> bool {
3486    if trusted {
3487        return true;
3488    }
3489    match GRAPH_RACE_STATE.load(Ordering::Relaxed) {
3490        1 => true,
3491        2 => false,
3492        _ => GRAPH_RACE_ARM_GRAPH.load(Ordering::Relaxed) == 1,
3493    }
3494}
3495
3496/// First decode token of a racing graph generation: hopeless already?
3497/// (>4x the normal path's per-token average AND over a second.) Settles
3498/// the race immediately; the caller discards the graph result and
3499/// recomputes this token on the normal path.
3500pub fn graph_race_first_token_hopeless(dur: std::time::Duration) -> bool {
3501    if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
3502        return false;
3503    }
3504    let first = GRAPH_RACE_TOK.load(Ordering::Relaxed) == 0;
3505    let cn = GRAPH_N[0].load(Ordering::Relaxed);
3506    if !first || cn == 0 {
3507        return false;
3508    }
3509    let c_avg = GRAPH_NS[0].load(Ordering::Relaxed) / cn as u64;
3510    let ns = dur.as_nanos() as u64;
3511    if ns > 1_000_000_000 && ns > 4 * c_avg {
3512        GRAPH_RACE_STATE.store(2, Ordering::Relaxed);
3513        tracing::info!(
3514            "wgpu graph race: first graph token {:.0} ms vs normal {:.2} ms/tok — hopeless, normal path wins",
3515            ns as f64 / 1e6,
3516            c_avg as f64 / 1e6
3517        );
3518        return true;
3519    }
3520    false
3521}
3522
3523/// Record one decode-token wall time for the racing arm. The first
3524/// token of each generation is discarded (KV-mirror upload / cold
3525/// caches on the graph arm; cold mmap on the normal arm).
3526pub fn graph_race_record(used_graph: bool, dur: std::time::Duration) {
3527    if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
3528        return;
3529    }
3530    let tok = GRAPH_RACE_TOK.fetch_add(1, Ordering::Relaxed);
3531    if tok == 0 {
3532        return;
3533    }
3534    let i = used_graph as usize;
3535    GRAPH_NS[i].fetch_add(dur.as_nanos() as u64, Ordering::Relaxed);
3536    GRAPH_N[i].fetch_add(1, Ordering::Relaxed);
3537}
3538
3539/// Bounded-cost content fingerprint for the backends' pointer-keyed device
3540/// caches: FNV over the whole slice up to 4 KiB, over 64 spread 64-byte
3541/// windows (plus the length) above. An address-keyed hit must also prove
3542/// the bytes are still the ones it uploaded — the allocator reuses heap
3543/// and mmap addresses freely, so a reloaded model or a re-dequantized
3544/// layer lands where the old bytes were — and sampling keeps that proof at
3545/// ~a microsecond even for a 126 MB matrix. Real replacements (another
3546/// model's tensor, an Adam-updated master) differ densely, so a 4 KiB
3547/// spread cannot miss them.
3548pub(crate) fn fp_bytes(data: &[u8]) -> u64 {
3549    #[inline]
3550    fn fnv(mut h: u64, bytes: &[u8]) -> u64 {
3551        let (chunks, tail) = bytes.split_at(bytes.len() & !7);
3552        for c in chunks.chunks_exact(8) {
3553            h ^= u64::from_le_bytes(c.try_into().unwrap());
3554            h = h.wrapping_mul(0x100_0000_01b3);
3555        }
3556        for &b in tail {
3557            h ^= b as u64;
3558            h = h.wrapping_mul(0x100_0000_01b3);
3559        }
3560        h
3561    }
3562    let mut h = 0xcbf2_9ce4_8422_2325u64 ^ (data.len() as u64);
3563    if data.len() <= 4096 {
3564        return fnv(h, data);
3565    }
3566    let step = (data.len() - 64) / 63;
3567    for i in 0..64 {
3568        h = fnv(h, &data[i * step..i * step + 64]);
3569    }
3570    h
3571}
3572
3573/// `fp_bytes` over an f32 slice without a bytemuck dependency (the Metal
3574/// backend builds with no GPU feature flags).
3575pub(crate) fn fp_f32(data: &[f32]) -> u64 {
3576    let bytes = unsafe { std::slice::from_raw_parts(data.as_ptr() as *const u8, data.len() * 4) };
3577    fp_bytes(bytes)
3578}
3579
3580#[cfg(test)]
3581mod fp_tests {
3582    use super::fp_bytes;
3583
3584    /// The pointer-keyed caches survive on `fp_bytes` telling two different
3585    /// tensors apart at a reused address. Its sampling must therefore see a
3586    /// change ANYWHERE — head, tail, and the stretches between windows are
3587    /// the places a cheaper hash would go blind.
3588    #[test]
3589    fn fp_bytes_sees_a_change_anywhere_in_a_sampled_slice() {
3590        let n = 1 << 20; // 1 MiB — far above the 4 KiB full-hash threshold
3591        let base: Vec<u8> = (0..n).map(|i| (i * 31 + 7) as u8).collect();
3592        let h0 = fp_bytes(&base);
3593        assert_eq!(h0, fp_bytes(&base), "fingerprint must be deterministic");
3594        // A DENSE change (every requantized/redequantized tensor is one)
3595        // must flip the fingerprint no matter how the windows fall.
3596        let mut dense = base.clone();
3597        for b in dense.iter_mut() {
3598            *b = b.wrapping_add(1);
3599        }
3600        assert_ne!(
3601            h0,
3602            fp_bytes(&dense),
3603            "a fully different tensor slipped through"
3604        );
3605        // Length participates: the same prefix at a shorter length is a
3606        // different key AND a different fingerprint.
3607        assert_ne!(h0, fp_bytes(&base[..n - 64]));
3608        // Below the threshold the hash is exact: a single flipped byte in
3609        // a norm-sized vector must be seen.
3610        let mut small = vec![3u8; 4096];
3611        let hs = fp_bytes(&small);
3612        small[2048] ^= 1;
3613        assert_ne!(hs, fp_bytes(&small), "full hash missed a one-byte change");
3614        // And the sampled windows land within bounds on awkward sizes.
3615        for n in [4097usize, 5000, 64 * 64, 1 << 16] {
3616            let v = vec![9u8; n];
3617            let _ = fp_bytes(&v); // must not panic on window math
3618        }
3619    }
3620}
3621
3622/// Hand the card back after a bake: drop its resident weights, planes and
3623/// pools so the ordinary engine (the runtime gate, a serve that follows)
3624/// starts from a clean budget. No-op off the wgpu backend.
3625pub fn bake_release() {
3626    #[cfg(feature = "gpu")]
3627    crate::gpu_wgpu::bake_release();
3628}
3629
3630/// Strict-f32 for the bake's GEMMs (phase A mask training): the mask
3631/// selects neurons by a gradient signal, and f16 operand rounding on
3632/// that signal closes the wrong ones. No-op off the wgpu backend.
3633pub fn bake_precision_strict(on: bool) {
3634    #[cfg(feature = "gpu")]
3635    crate::gpu_wgpu::bake_precision_strict(on);
3636    #[cfg(not(feature = "gpu"))]
3637    let _ = on;
3638}
3639
3640/// CMF_GRAPH_HOSTPROF=1: how a graph token's wall splits between the
3641/// host encoding the command stream and the tail the GPU still owes
3642/// after encode. Fifteen GPU-side suspects measured null while the
3643/// bench counted 17.7k allocations a token — this is the instrument
3644/// that says whether the thief was on the host all along.
3645pub fn hostprof_encode_done(t0: std::time::Instant) {
3646    use std::sync::atomic::{AtomicU64, Ordering};
3647    static ENC: AtomicU64 = AtomicU64::new(0);
3648    static N: AtomicU64 = AtomicU64::new(0);
3649    if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
3650        return;
3651    }
3652    ENC.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
3653    let n = N.fetch_add(1, Ordering::Relaxed) + 1;
3654    if n % 100 == 0 {
3655        eprintln!(
3656            "hostprof: encode {:.2} ms/token over {n} tokens",
3657            ENC.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
3658        );
3659    }
3660}
3661
3662pub fn hostprof_total(t0: std::time::Instant) {
3663    use std::sync::atomic::{AtomicU64, Ordering};
3664    static TOT: AtomicU64 = AtomicU64::new(0);
3665    static N: AtomicU64 = AtomicU64::new(0);
3666    if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
3667        return;
3668    }
3669    TOT.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
3670    let n = N.fetch_add(1, Ordering::Relaxed) + 1;
3671    if n % 100 == 0 {
3672        eprintln!(
3673            "hostprof: total {:.2} ms/token over {n} tokens",
3674            TOT.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
3675        );
3676    }
3677}
3678
3679/// Per-stage host-encode accumulator for the Metal token loop
3680/// (CMF_GRAPH_HOSTPROF=1). Stage 0 = GDN-run encode; everything else
3681/// falls out by subtraction from hostprof's encode total.
3682pub fn stageprof(stage: u32, dt: std::time::Duration) {
3683    use std::sync::atomic::{AtomicU64, Ordering};
3684    static NS: [AtomicU64; 4] = [
3685        AtomicU64::new(0),
3686        AtomicU64::new(0),
3687        AtomicU64::new(0),
3688        AtomicU64::new(0),
3689    ];
3690    static N: AtomicU64 = AtomicU64::new(0);
3691    if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
3692        return;
3693    }
3694    NS[stage as usize % 4].fetch_add(dt.as_nanos() as u64, Ordering::Relaxed);
3695    if stage == 1 {
3696        let n = N.fetch_add(1, Ordering::Relaxed) + 1;
3697        if n % 200 == 0 {
3698            eprintln!(
3699                "stageprof: planning {:.2} ms/tok | gdn-item {:.2} ms/tok | attn-item {:.2} ms/tok ({n} tok)",
3700                NS[1].load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
3701                NS[2].load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
3702                NS[3].load(Ordering::Relaxed) as f64 / n as f64 / 1e6
3703            );
3704        }
3705    }
3706}
3707
3708/// Active weight bytes dispatched so far (Metal decode path); 0 where
3709/// the backend does not count. The honest floor's numerator.
3710pub fn weight_bytes_dispatched() -> u64 {
3711    let mut total = 0u64;
3712    #[cfg(target_os = "macos")]
3713    {
3714        total += crate::gpu_metal::WEIGHT_BYTES.load(std::sync::atomic::Ordering::Relaxed);
3715    }
3716    #[cfg(feature = "gpu")]
3717    {
3718        total += crate::gpu_wgpu::WEIGHT_BYTES.load(std::sync::atomic::Ordering::Relaxed);
3719    }
3720    total
3721}
3722
3723/// The per-stage split of `weight_bytes_dispatched`:
3724/// [misc, dense-ffn, moe, attn, gdn, head].
3725pub fn weight_bytes_by() -> [u64; 6] {
3726    #[cfg(target_os = "macos")]
3727    {
3728        let mut o = [0u64; 6];
3729        for (i, a) in crate::gpu_metal::WEIGHT_BYTES_BY.iter().enumerate() {
3730            o[i] = a.load(std::sync::atomic::Ordering::Relaxed);
3731        }
3732        return o;
3733    }
3734    #[allow(unreachable_code)]
3735    [0; 6]
3736}
3737
3738#[cfg(test)]
3739mod probe_warmup_tests {
3740    use super::*;
3741    use std::time::Duration;
3742
3743    fn ms(v: f64) -> Duration {
3744        Duration::from_nanos((v * 1e6) as u64)
3745    }
3746
3747    /// The bug this pins, measured on an A100: the first device call for
3748    /// a class compiles its pipeline, was timed at 117.01 ms against the
3749    /// host's 3.19, and sent `gemm-nt` to the CPU for the whole process —
3750    /// which ran a 27B bake on 2.6 cores with the card idle.
3751    #[test]
3752    fn one_cold_first_sample_does_not_lose_the_class() {
3753        let p = Probe::new();
3754        // First device sample is the pipeline build. Then the truth.
3755        probe_record_into(&p, "gemm-nt", None, true, ms(117.01));
3756        probe_record_into(&p, "gemm-nt", None, true, ms(1.1));
3757        probe_record_into(&p, "gemm-nt", None, true, ms(1.0));
3758        probe_record_into(&p, "gemm-nt", None, false, ms(3.19));
3759        probe_record_into(&p, "gemm-nt", None, false, ms(3.20));
3760        assert_eq!(
3761            p.state.load(Ordering::Relaxed),
3762            1,
3763            "the device is 3x faster once warm and must win"
3764        );
3765    }
3766
3767    /// The warm-up must not become a way to never decide, and must not
3768    /// underflow: a blind decrement at zero wraps a u32 to its maximum
3769    /// and mutes the arm for the life of the process.
3770    #[test]
3771    fn the_warmup_is_spent_once_and_never_underflows() {
3772        let p = Probe::new();
3773        for _ in 0..8 {
3774            probe_record_into(&p, "matmat", None, true, ms(10.0));
3775        }
3776        assert_eq!(p.gpu_burn.load(Ordering::Relaxed), 0, "spent, not wrapped");
3777        assert_eq!(
3778            p.gpu_n.load(Ordering::Relaxed),
3779            7,
3780            "one sample burned, the rest counted"
3781        );
3782    }
3783
3784    /// A device path that always refuses records no timing, so without
3785    /// counting the refusals the class can never reach a verdict. On an
3786    /// M4 with LFM2.5-2.6B `ffn` was still undecided after 9000 calls,
3787    /// alternating arms and paying a failed device attempt on half of
3788    /// them.
3789    #[test]
3790    fn a_class_whose_device_always_declines_settles_on_the_host() {
3791        let _probe_guard = probe_test_guard();
3792        // A class no other test in this file touches: `probe_note_decline`
3793        // works on the process-wide probes by design, and the tests in
3794        // this binary share them.
3795        let c = OpClass::MatmatWide;
3796        let p = &PROBES[c as usize];
3797        p.state.store(0, Ordering::Relaxed);
3798        p.declines.store(0, Ordering::Relaxed);
3799        for _ in 0..(PROBE_DECLINE_LIMIT - 1) {
3800            probe_note_decline(c);
3801        }
3802        assert_eq!(
3803            p.state.load(Ordering::Relaxed),
3804            0,
3805            "one short of the limit is still a question, not an answer"
3806        );
3807        probe_note_decline(c);
3808        assert_eq!(p.state.load(Ordering::Relaxed), 2, "settled on the host");
3809        assert!(matches!(probe_arm(c), ProbeArm::Cpu));
3810        p.state.store(0, Ordering::Relaxed);
3811        p.declines.store(0, Ordering::Relaxed);
3812    }
3813
3814    /// A genuinely slower device still loses — the warm-up removes an
3815    /// artefact, it does not put a thumb on the scale.
3816    #[test]
3817    fn a_slow_device_still_loses_after_the_warmup() {
3818        let p = Probe::new();
3819        for _ in 0..4 {
3820            probe_record_into(&p, "matvec", None, true, ms(40.0));
3821        }
3822        for _ in 0..4 {
3823            probe_record_into(&p, "matvec", None, false, ms(2.0));
3824        }
3825        assert_eq!(p.state.load(Ordering::Relaxed), 2, "host wins on merit");
3826    }
3827}
3828
3829/// Scratch/weight lifetime for a synchronous image-pipeline stage. Declare
3830/// this before the stage model so the model drops before cache collection.
3831pub(crate) struct ImageStageGuard {
3832    #[cfg(target_os = "macos")]
3833    metal: Option<crate::gpu_metal::ImageStageGuard>,
3834    #[cfg(feature = "gpu")]
3835    wgpu: crate::gpu_wgpu::ImageStageGuard,
3836}
3837
3838pub(crate) fn image_stage_scope() -> ImageStageGuard {
3839    ImageStageGuard {
3840        #[cfg(target_os = "macos")]
3841        metal: if matches!(backend(), Backend::Metal) {
3842            Some(crate::gpu_metal::image_stage_scope())
3843        } else {
3844            None
3845        },
3846        #[cfg(feature = "gpu")]
3847        wgpu: crate::gpu_wgpu::image_stage_scope(),
3848    }
3849}
3850
3851impl ImageStageGuard {
3852    pub(crate) fn track_model(&mut self, uid: u64) {
3853        #[cfg(target_os = "macos")]
3854        if let Some(metal) = &mut self.metal {
3855            metal.track_model(uid);
3856        }
3857        #[cfg(feature = "gpu")]
3858        self.wgpu.track_model(uid);
3859        #[cfg(not(target_os = "macos"))]
3860        let _ = uid;
3861    }
3862}
3863
3864// ════════════════════════════════════════════════════════════════════
3865// Z-Image-Turbo device contract (WP0 scaffold, plan §2.1). APPEND-ONLY.
3866//
3867// Owner of the contract: the WP1 lead. The backends implement it in their
3868// own child modules — `gpu_wgpu/zimage.rs` (WP2) and `gpu_metal/zimage.rs`
3869// (WP3) — and never edit the parent files. New fields are added only as
3870// `Option<…>` with agreed semantics; existing fields never change meaning.
3871//
3872// Convention (the same as every `gpu::*` entry): `false` = "not handled",
3873// nothing observable was changed, and the caller runs the CPU path
3874// (`zimage::ZImageDit::step_cpu` etc.), which is the bit-level reference.
3875//
3876// Sequence order everywhere is diffusers' [img rows…, cap rows…], with
3877// padded lengths n_img_p = ceil32(n_img) and n_cap_p = ceil32(L).
3878// ════════════════════════════════════════════════════════════════════
3879
3880/// One Z-Image transformer block's device inputs (noise refiner, context
3881/// refiner or main layer — all share this shape). Weights are tensor
3882/// indices into `model.tensors` (diffusers names under `dit.`); the codec
3883/// is whatever the container holds (F16/Bf16/Q8Row/Q8_2f/Q4TiledP…), and a
3884/// backend that cannot expand a codec declines (returns `false`).
3885/// Norm vectors are f32 host slices that live as long as the caller's
3886/// `ZImageDit`; a backend may cache them by pointer (they do not change).
3887#[derive(Clone, Copy)]
3888pub struct ZBlockRef<'a> {
3889    /// `attention.to_q/to_k/to_v/to_out.0.weight`, each [hidden, hidden].
3890    pub wq: usize,
3891    pub wk: usize,
3892    pub wv: usize,
3893    pub wo: usize,
3894    /// `feed_forward.w1` (gate) / `w3` (up) [inter, hidden], `w2` (down)
3895    /// [hidden, inter]. FFN = w2(silu(w1·x) ⊙ w3·x).
3896    pub w1: usize,
3897    pub w3: usize,
3898    pub w2: usize,
3899    /// `attention_norm1` / `attention_norm2`, [hidden] (plain-w RMSNorm).
3900    pub norm1: &'a [f32],
3901    pub norm2: &'a [f32],
3902    /// `ffn_norm1` / `ffn_norm2`, [hidden].
3903    pub ffn_norm1: &'a [f32],
3904    pub ffn_norm2: &'a [f32],
3905    /// `attention.norm_q` / `norm_k`, [hd] (per-head RMSNorm before RoPE).
3906    pub norm_q: &'a [f32],
3907    pub norm_k: &'a [f32],
3908}
3909
3910/// Z-Image geometry. Turbo: hidden 3840, nh 30 (MHA, no GQA), hd 128,
3911/// inter 10240, eps 1e-5 (all RMSNorms incl. qk-norm), final_eps 1e-6
3912/// (the affine-free final LayerNorm), patch_dim 64 (2×2×16).
3913#[derive(Clone, Copy, Debug, PartialEq)]
3914pub struct ZGeom {
3915    pub hidden: usize,
3916    pub nh: usize,
3917    pub hd: usize,
3918    pub inter: usize,
3919    pub eps: f32,
3920    pub final_eps: f32,
3921    pub patch_dim: usize,
3922}
3923
3924/// Once per (prompt, resolution). The backend uploads/caches what it needs
3925/// keyed by `key`; weight planes are keyed by the MODEL (not by `key`) and
3926/// survive across prompts until `zimage_release`.
3927pub struct ZPrepareArgs<'a> {
3928    pub model: &'a Arc<CmfModel>,
3929    pub geom: ZGeom,
3930    /// Caller-chosen identity of this (prompt, resolution) state; every
3931    /// `ZStepArgs` of the same image carries the same key.
3932    pub key: u64,
3933    /// Image tokens (H/16 · W/16), padded count ceil32(n_img), caption
3934    /// padded count ceil32(L). S = n_img_p + n_cap_p.
3935    pub n_img: usize,
3936    pub n_img_p: usize,
3937    pub n_cap_p: usize,
3938    /// The patch grid (H/16, W/16); n_img = grid.0 · grid.1. Row-major
3939    /// token order `hp·grid.1 + wp`.
3940    pub grid: (usize, usize),
3941    /// [n_cap_p, hidden], ALREADY context-refined (host or device).
3942    pub cap: &'a [f32],
3943    /// Noise-refiner RoPE: [n_img_p · hd/2] cos, sin (complex-interleaved
3944    /// pairs, hd/2 angles per token).
3945    pub rope_img: (&'a [f32], &'a [f32]),
3946    /// Main-layer RoPE: [(n_img_p + n_cap_p) · hd/2], rows ordered [img, cap].
3947    pub rope_joint: (&'a [f32], &'a [f32]),
3948    /// `all_x_embedder.2-1.weight` [hidden, 64], `.bias` [hidden],
3949    /// `x_pad_token` [hidden] (replaces rows ≥ n_img after the embed).
3950    pub x_emb_w: &'a [f32],
3951    pub x_emb_b: &'a [f32],
3952    pub x_pad: &'a [f32],
3953    /// `all_final_layer.2-1.linear.weight` [64, hidden], `.bias` [64].
3954    pub final_w: &'a [f32],
3955    pub final_b: &'a [f32],
3956    /// 2 noise-refiner blocks (image rows only) and 30 main layers.
3957    pub noise_refiner: &'a [ZBlockRef<'a>],
3958    pub layers: &'a [ZBlockRef<'a>],
3959    /// OPTIONAL (backends may ignore): the modulation of EVERY step of this
3960    /// image, [steps][(2+30)·4·hidden] in the `ZStepArgs::mods` layout, and
3961    /// [steps][hidden] final scales, so a backend can upload them once per
3962    /// image and index them by `ZStepArgs::step`. `ZStepArgs::mods` is still
3963    /// always supplied and is authoritative.
3964    pub mods_all: Option<&'a [f32]>,
3965    pub final_scale_all: Option<&'a [f32]>,
3966    /// OPTIONAL (B2): the CFG negative item. When `Some`, the backend
3967    /// prepares ONE batch-2 program under `key` — item 0 is this prompt,
3968    /// item 1 the negative — and every `ZStepArgs` of that key must carry
3969    /// `out_neg`. A backend without batch 2 returns `false` (the caller
3970    /// then prepares the two items separately or runs the CPU path).
3971    pub neg: Option<ZNegArgs<'a>>,
3972}
3973
3974/// The negative (unconditional) item of a CFG pair: its own refined
3975/// caption, padded caption length and joint RoPE table (the image ids sit
3976/// at axis-0 position L_p+1, so both tables depend on the item's L_p).
3977pub struct ZNegArgs<'a> {
3978    /// [n_cap_p, hidden], context-refined.
3979    pub cap: &'a [f32],
3980    pub n_cap_p: usize,
3981    /// [n_img_p · hd/2] cos, sin (noise refiner) of this item.
3982    pub rope_img: (&'a [f32], &'a [f32]),
3983    /// [(n_img_p + n_cap_p) · hd/2] cos, sin, rows [img, cap].
3984    pub rope_joint: (&'a [f32], &'a [f32]),
3985}
3986
3987/// Once per denoising step.
3988pub struct ZStepArgs<'a> {
3989    /// The `ZPrepareArgs::key` this step belongs to. A key the backend has
3990    /// not prepared → `false`.
3991    pub key: u64,
3992    /// Step index into the schedule (0..steps); selects the row of
3993    /// `ZPrepareArgs::mods_all` when a backend uses it.
3994    pub step: usize,
3995    /// [n_img_p, 64] patchified latent, inner order (dy·2+dx)·16+c. Rows
3996    /// ≥ n_img are copies of the last row; the backend replaces them with
3997    /// `x_pad` after the embed.
3998    pub x_tok: &'a [f32],
3999    /// Per block (noise_refiner then layers) the RAW chunks
4000    /// [scale_msa, gate_msa, scale_mlp, gate_mlp] of Linear(temb) (no SiLU
4001    /// before it), [(2+30)·4·hidden]. The backend applies (1+s) and tanh(g).
4002    pub mods: &'a [f32],
4003    /// [hidden] = 1 + Linear(SiLU(temb)) — already includes the +1.
4004    pub final_scale: &'a [f32],
4005    /// [n_img, 64]: the model output v (before the pipeline's negation),
4006    /// image rows only, patchified order.
4007    pub out: &'a mut [f32],
4008    /// [n_img, 64]: the negative item's v — required (and only valid) for
4009    /// a key prepared with `ZPrepareArgs::neg`. Both items see `x_tok`.
4010    pub out_neg: Option<&'a mut [f32]>,
4011}
4012
4013/// Prepare the per-(prompt, resolution) device state. Backends: wgpu →
4014/// `gpu_wgpu::zimage::prepare` (WP2), Metal → `gpu_metal::zimage::prepare`
4015/// (WP3).
4016#[allow(unused_variables)]
4017pub fn zimage_prepare(a: &ZPrepareArgs) -> bool {
4018    match backend() {
4019        #[cfg(target_os = "macos")]
4020        Backend::Metal => crate::gpu_metal::zimage::prepare(a),
4021        #[cfg(feature = "gpu")]
4022        Backend::Wgpu => crate::gpu_wgpu::zimage::prepare(a),
4023        #[allow(unreachable_patterns)]
4024        _ => false,
4025    }
4026}
4027
4028/// One full DiT forward on the device: x_embed → pad rows → noise refiner
4029/// ×2 → concat [img, cap] → 30 layers → final LayerNorm·scale → Linear →
4030/// image rows into `a.out`.
4031#[allow(unused_variables)]
4032pub fn zimage_step(a: &mut ZStepArgs) -> bool {
4033    match backend() {
4034        #[cfg(target_os = "macos")]
4035        Backend::Metal => crate::gpu_metal::zimage::step(a),
4036        #[cfg(feature = "gpu")]
4037        Backend::Wgpu => crate::gpu_wgpu::zimage::step(a),
4038        #[allow(unreachable_patterns)]
4039        _ => false,
4040    }
4041}
4042
4043/// OPTIONAL (B2): build the backend's weight planes for the per-step blocks
4044/// and the context refiner ahead of `zimage_prepare`, so the caller can
4045/// overlap the upload with the (CPU) text encoder. `false` = not done;
4046/// `zimage_prepare` builds whatever is missing either way.
4047#[allow(unused_variables)]
4048pub fn zimage_preload(
4049    model: &Arc<CmfModel>,
4050    geom: &ZGeom,
4051    noise_refiner: &[ZBlockRef],
4052    layers: &[ZBlockRef],
4053    context_refiner: &[ZBlockRef],
4054) -> bool {
4055    match backend() {
4056        #[cfg(target_os = "macos")]
4057        Backend::Metal => {
4058            crate::gpu_metal::zimage::preload(model, geom, noise_refiner, layers, context_refiner)
4059        }
4060        #[cfg(feature = "gpu")]
4061        Backend::Wgpu => crate::gpu_wgpu::zimage::preload(model, geom, noise_refiner, layers, context_refiner),
4062        #[allow(unreachable_patterns)]
4063        _ => false,
4064    }
4065}
4066
4067/// Persist the driver's compiled pipelines after a Z-Image generation (the
4068/// chain's kernels are built at first use, after the context came up), so
4069/// the next process skips the compile. Best-effort, no-op off wgpu.
4070pub fn zimage_flush_pipelines() {
4071    #[cfg(feature = "gpu")]
4072    if matches!(backend(), Backend::Wgpu) {
4073        crate::gpu_wgpu::pipeline_cache_flush();
4074    }
4075}
4076
4077/// OPTIONAL (B2): bring the device up and compile the Z-Image kernels, on
4078/// a helper thread at the start of a generation (the context and the
4079/// compiles cost ~1 s cold, beside the host-side loading). `false` = no
4080/// device path here.
4081pub fn zimage_warmup() -> bool {
4082    match backend() {
4083        #[cfg(target_os = "macos")]
4084        Backend::Metal => crate::gpu_metal::zimage::warmup(),
4085        #[cfg(feature = "gpu")]
4086        Backend::Wgpu => crate::gpu_wgpu::zimage::warmup(),
4087        #[allow(unreachable_patterns)]
4088        _ => false,
4089    }
4090}
4091
4092/// OPTIONAL (B2): upload the resident VAE's weights and compile its
4093/// kernels ahead of `vae_decode_chain` (the caller runs it on a helper
4094/// thread while the DiT steps keep the device busy). `false` = not done.
4095#[allow(unused_variables)]
4096pub fn vae_prewarm(a: &crate::vae::VaeChainArgs) -> bool {
4097    match backend() {
4098        #[cfg(target_os = "macos")]
4099        Backend::Metal => crate::gpu_metal::zimage::vae_prewarm(a),
4100        #[cfg(feature = "gpu")]
4101        Backend::Wgpu => crate::gpu_wgpu::zimage::vae_prewarm(a),
4102        #[allow(unreachable_patterns)]
4103        _ => false,
4104    }
4105}
4106
4107/// Drop the Z-Image DiT device state (planes, prepared programs) but keep
4108/// the VAE chain (B2: the generator frees the DiT before decoding).
4109pub fn zimage_release_dit() {
4110    #[cfg(target_os = "macos")]
4111    crate::gpu_metal::zimage::release_dit();
4112    #[cfg(feature = "gpu")]
4113    crate::gpu_wgpu::zimage::release_dit();
4114}
4115
4116/// Drop every Z-Image device resource (planes, prepared states, VAE chain
4117/// buffers): stage change or process end. Calls each compiled backend's
4118/// release directly, without `backend()`, so it never brings a device up;
4119/// the child modules' `release` must touch module-local state only.
4120pub fn zimage_release() {
4121    #[cfg(target_os = "macos")]
4122    crate::gpu_metal::zimage::release();
4123    #[cfg(feature = "gpu")]
4124    crate::gpu_wgpu::zimage::release();
4125}
4126
4127/// Optional device context refiner: the same block math with scale = 0 and
4128/// gate = 1 (unmodulated): x += norm2(attn(norm1(x))); x += ffn_norm2(ffn(
4129/// ffn_norm1(x))). `cap` is [n_cap_p, hidden] in/out (the cap_embedder
4130/// output with pad rows already = cap_pad_token); `rope_cap` is
4131/// [n_cap_p · hd/2] cos, sin. `false` = untouched, run the CPU refiner.
4132#[allow(unused_variables)]
4133pub fn zimage_refine_caption(
4134    model: &Arc<CmfModel>,
4135    geom: &ZGeom,
4136    blocks: &[ZBlockRef],
4137    rope_cap: (&[f32], &[f32]),
4138    cap: &mut [f32],
4139) -> bool {
4140    match backend() {
4141        #[cfg(target_os = "macos")]
4142        Backend::Metal => {
4143            crate::gpu_metal::zimage::refine_caption(model, geom, blocks, rope_cap, cap)
4144        }
4145        #[cfg(feature = "gpu")]
4146        Backend::Wgpu => {
4147            crate::gpu_wgpu::zimage::refine_caption(model, geom, blocks, rope_cap, cap)
4148        }
4149        #[allow(unreachable_patterns)]
4150        _ => false,
4151    }
4152}
4153
4154/// Resident Flux-VAE decode (the whole decoder on the device, one latent
4155/// upload, one RGB readback). `a` comes from `VaeDecoder::chain_args()`.
4156/// `z` is [latent_channels, h, w] ALREADY de-normalised
4157/// (z/scaling_factor + shift_factor — the conv_in input); `out` is
4158/// [3, 8h, 8w], the raw decoder output (≈[-1, 1], before x/2+0.5).
4159#[allow(unused_variables)]
4160pub fn vae_decode_chain(
4161    a: &crate::vae::VaeChainArgs,
4162    z: &[f32],
4163    h: usize,
4164    w: usize,
4165    out: &mut [f32],
4166) -> bool {
4167    match backend() {
4168        #[cfg(target_os = "macos")]
4169        Backend::Metal => crate::gpu_metal::zimage::vae_decode_chain(a, z, h, w, out),
4170        #[cfg(feature = "gpu")]
4171        Backend::Wgpu => crate::gpu_wgpu::zimage::vae_decode_chain(a, z, h, w, out),
4172        #[allow(unreachable_patterns)]
4173        _ => false,
4174    }
4175}