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        /// Absolute position of `cpu_k`/`cpu_v` row 0
1571        /// (`LayerKvCache::base`): a sliding layer that keeps only its
1572        /// tail stores rows from there on. 0 for every untrimmed layer.
1573        cpu_base: usize,
1574        /// This layer's own attention geometry, when the model's layers do
1575        /// not share one (MiMo-V2: 4/8 KV heads, 128-wide V under 192-wide
1576        /// heads, sliding windows with learned sinks, two RoPE tables).
1577        /// None = the call-wide (nkv, hd, rd, invf), V as wide as K, full
1578        /// context and a plain softmax — the historical contract, whose
1579        /// kernels and dispatch are untouched.
1580        geom: Option<GraphAttnGeom<'a>>,
1581        /// Spark-X2.5 head-wise output gate: `self_attn.g_proj`
1582        /// `[num_heads, hidden]` against the attention's normed input; head
1583        /// h's output is scaled by sigmoid(g[h]) before the O projection.
1584        /// None = no such gate (every other model).
1585        head_gate: Option<GraphW<'a>>,
1586    },
1587    Gdn {
1588        qkv: GraphW<'a>,
1589        z: GraphW<'a>,
1590        a: GraphW<'a>,
1591        b: GraphW<'a>,
1592        out: GraphW<'a>,
1593        conv1d: &'a [f32],
1594        a_log: &'a [f32],
1595        dt_bias: &'a [f32],
1596        norm: &'a [f32],
1597        nv: usize,
1598        nk: usize,
1599        dk: usize,
1600        dv: usize,
1601        kk: usize,
1602        /// CPU recurrent state `[ring (kk-1)·cdim | S nv·dk·dv]` — seeds the
1603        /// device mirror when prefill ran on the host (o1 collection, CPU
1604        /// fallback): a zero-initialized device state at decode is exactly
1605        /// the "coherent but contextless" garble.
1606        cpu_state: &'a [f32],
1607    },
1608    /// LFM2 gated short convolution: a fused (B, C, x) projection, a
1609    /// depthwise causal conv over a (kernel−1)-deep per-channel ring,
1610    /// C-gating, and an output projection. This mixer is what most of an
1611    /// LFM2 stack is (22 of the 2.6B's 30 layers), and before it had a
1612    /// graph arm the whole model fell to the per-op path — ~100 submits
1613    /// a token, 22 tok/s on an A100 for a 1.4 GB file.
1614    ShortConv {
1615        /// [3·hidden, hidden] fused input projection.
1616        inp: GraphW<'a>,
1617        /// [hidden, hidden] output projection.
1618        out: GraphW<'a>,
1619        /// [hidden · kernel] depthwise taps, `[channel][tap]`, tap
1620        /// kernel−1 multiplying the current position.
1621        taps: &'a [f32],
1622        kernel: usize,
1623        /// CPU conv ring `[channel][kernel−1]`, slot 0 newest — seeds
1624        /// the device mirror when prefill ran on the host, which for
1625        /// this mixer is always (the batch graph declines it).
1626        cpu_state: &'a [f32],
1627    },
1628}
1629
1630/// Per-layer attention geometry for the wgpu graphs (see
1631/// `GraphAttn::Full::geom`). The layer's CPU cache keeps K rows `hd` wide
1632/// and V rows zero-padded to `hd`; the device mirror stores V `dv` wide,
1633/// and a windowed layer keeps a ring of the last positions only.
1634#[derive(Clone, Copy)]
1635pub struct GraphAttnGeom<'a> {
1636    /// KV heads of this layer (divides the Q heads).
1637    pub nkv: usize,
1638    /// V head width, `4 <= dv <= head_dim`, a multiple of 4.
1639    pub dv: usize,
1640    /// Rotary width of this layer (NeoX half-split over `[0, rd)`).
1641    pub rd: usize,
1642    /// This layer's RoPE inverse frequencies (`rd / 2` of them).
1643    pub invf: &'a [f32],
1644    /// Positions a query sees, its own included (MiMo-V2 SWA: 128);
1645    /// None = the whole context.
1646    pub window: Option<usize>,
1647    /// Learned per-Q-head sink logits (gpt-oss / MiMo-V2): they join the
1648    /// softmax max and denominator and carry no value row.
1649    pub sink: Option<&'a [f32]>,
1650}
1651
1652/// Activation of a graph layer's dense FFN, `act(gate)·up`. The code is
1653/// the word the kernels read (`silu_mul_pre`'s `_c`, the fused gate+up
1654/// kernels' `_p1`): 0 keeps the historical SiLU expression untouched.
1655#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1656pub enum GraphAct {
1657    Silu,
1658    /// Exact erf GELU (HF `gelu`), erf by Abramowitz–Stegun 7.1.26 — the
1659    /// formula `inference::erf_f32` uses on the host.
1660    GeluErf,
1661}
1662
1663impl GraphAct {
1664    pub fn code(self) -> u32 {
1665        match self {
1666            Self::Silu => 0,
1667            Self::GeluErf => 1,
1668        }
1669    }
1670}
1671
1672/// Per-layer weights for the whole-token wgpu graph.
1673pub struct GraphLayer<'a> {
1674    pub input_norm: &'a [f32],
1675    pub attn: GraphAttn<'a>,
1676    pub post_norm: &'a [f32],
1677    pub ffn: GraphFfn<'a>,
1678}
1679
1680/// The FFN of one graph layer: a dense SwiGLU trio, or a routed MoE —
1681/// router + top-k selection + all selected experts run ON DEVICE (the
1682/// routing decision depends on the resident hidden state, so a CPU
1683/// round-trip per layer would forfeit the one-submit design).
1684pub enum GraphFfn<'a> {
1685    /// A singleton attention-only batch graph. Returns the post-attention
1686    /// residual, allowing a dynamic expert bank to own the FFN separately.
1687    AttentionOnly,
1688    Dense {
1689        gate: GraphW<'a>,
1690        up: GraphW<'a>,
1691        down: GraphW<'a>,
1692        /// act(gate)·up. The builders refuse an activation with no arm here
1693        /// (before this field the graph computed SiLU for any dense FFN).
1694        act: GraphAct,
1695    },
1696    Moe {
1697        /// Router logits weight (f32, kind 4) `[n_exp, hidden]`.
1698        router: GraphW<'a>,
1699        /// Shared-expert sigmoid gate (f32) `[1, hidden]`.
1700        shared_gate: GraphW<'a>,
1701        /// Per-expert q4_tiled directory indices `(gate, up, down)`;
1702        /// the SHARED expert rides as the LAST entry — the select
1703        /// kernel pins it with the sigmoid weight.
1704        experts: Vec<(usize, usize, usize)>,
1705        /// Routed experts (shared excluded).
1706        n_exp: usize,
1707        top_k: usize,
1708        inter: usize,
1709        norm_topk: bool,
1710        /// Expert weight layout, uniform across the layer: `false` =
1711        /// q4_tiled (18 B tiles, inline f16 scale), `true` = q4tp
1712        /// (16 B nibbles + a per-row ladder plane). The two differ only
1713        /// in where the scale comes from, so they share every kernel
1714        /// but the weight-staging block.
1715        q4tp: bool,
1716        /// `true` = the gate/up experts are `q2tp` (2-bit plane) while
1717        /// `down` stays q4tp — the mixed profile a 2-bit-class checkpoint
1718        /// converts into. Only meaningful with `q4tp: true`.
1719        gu_q2: bool,
1720        /// LFM2-MoE / DeepSeek-V3 `noaux_tc` routing: per-expert sigmoid
1721        /// scores instead of a softmax, and `norm_topk` renormalises with
1722        /// the 1e-6 floor. The softmax arm is bit-identical to before.
1723        sigmoid: bool,
1724        /// Per-expert SELECTION bias: added to the score for the top-k
1725        /// choice only — the mixing weights stay unbiased (noaux_tc).
1726        bias: Option<&'a [f32]>,
1727        /// Whether a shared expert rides as the last `experts` entry.
1728        /// LFM2-MoE has none; the select kernel then leaves slot `top_k`
1729        /// unwritten and the expert loop runs `top_k` slots, not +1.
1730        has_shared: bool,
1731        /// The shared expert carries a sigmoid gate (Qwen2/3-MoE). `false`
1732        /// with `has_shared`: the shared expert enters with weight 1
1733        /// (DeepSeek-V3 / HunYuan hy_v3) and `shared_gate` is a stand-in
1734        /// the select kernels ignore.
1735        shared_gated: bool,
1736        /// Multiplier on the routed mixing weights after the optional
1737        /// renormalization (`routed_scaling_factor`); 1.0 = none.
1738        route_scale: f32,
1739    },
1740}
1741
1742/// Outcome of one whole-token graph attempt. A failed attempt after sealed
1743/// O(1) state was admitted must not fall through to the stale CPU state.
1744#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1745pub enum TokenGraphOutcome {
1746    /// No command was committed; the caller may use its ordinary path.
1747    Declined,
1748    /// The graph completed and its hidden/logits output is valid.
1749    Completed,
1750    /// Sealed O(1) state was admitted and a later graph operation failed.
1751    Failed,
1752}
1753
1754/// Whole-token decode graph on wgpu: the entire layer stack in ONE submit,
1755/// hidden resident, one readback. Updates `h` in place.
1756/// `loop_norm_at`: virtual layer indices after which `final_norm` is applied
1757/// (Looped Transformer mid-stack norm). Empty for standard models.
1758#[allow(clippy::too_many_arguments)]
1759pub fn forward_token_graph(
1760    model: &Arc<CmfModel>,
1761    kv_id: u64,
1762    layers: &[GraphLayer],
1763    // Per-layer sealed o1 (Nystrom) state; Some = replace this layer's
1764    // exact attention with the O(1) kernels. wgpu only.
1765    o1: &[Option<Vec<crate::nystrom::O1DeviceView<'_>>>],
1766    o1_epoch: u64,
1767    invf: &[f32],
1768    h: &mut [f32],
1769    nh: usize,
1770    nkv: usize,
1771    hd: usize,
1772    attn_scale: f32,
1773    rd: usize,
1774    hidden: usize,
1775    inter: usize,
1776    position: usize,
1777    cap: usize,
1778    gemma: bool,
1779    eps: f32,
1780    lm_head: Option<(&GraphW, usize)>,
1781    final_norm: &[f32],
1782    logits: &mut Vec<f32>,
1783    loop_norm_at: &[usize],
1784    steps: usize,
1785    embed: Option<(&GraphW, usize, f32)>,
1786    ids_out: Option<&mut Vec<u32>>,
1787    // How many leading layers the graph ran (see the wgpu twin) — smaller
1788    // than layers.len() when the expert budget ended the device prefix.
1789    layers_run: Option<&mut usize>,
1790    // Absolute index of layers[0] in the model — the KV/state mirrors key
1791    // on it, so a layer SPAN (network split segment) shares mirrors with
1792    // a full-stack run instead of colliding at slot 0.
1793    layer_base: usize,
1794    // Read the final hidden back alongside the fused head's logits.
1795    hidden_too: bool,
1796) -> TokenGraphOutcome {
1797    match backend() {
1798        #[cfg(feature = "gpu")]
1799        Backend::Wgpu => crate::gpu_wgpu::forward_token_graph(
1800            model,
1801            kv_id,
1802            layers,
1803            o1,
1804            o1_epoch,
1805            invf,
1806            h,
1807            nh,
1808            nkv,
1809            hd,
1810            attn_scale,
1811            rd,
1812            hidden,
1813            inter,
1814            position,
1815            cap,
1816            gemma,
1817            eps,
1818            lm_head,
1819            final_norm,
1820            logits,
1821            loop_norm_at,
1822            steps,
1823            embed,
1824            ids_out,
1825            layers_run,
1826            layer_base,
1827            hidden_too,
1828        ),
1829        #[allow(unused_variables)]
1830        _ => {
1831            let _ = (
1832                attn_scale,
1833                lm_head,
1834                final_norm,
1835                logits,
1836                loop_norm_at,
1837                layers_run,
1838                layer_base,
1839                hidden_too,
1840            );
1841            TokenGraphOutcome::Declined
1842        }
1843    }
1844}
1845
1846/// Whole-token resident graph for the Embryo working model.  Unlike the
1847/// generic graph this path accepts a packed f32 model and keeps both phase
1848/// recurrent state and the anchor KV cache on the device.  `false` is an
1849/// honest capability refusal; callers must retain the CPU executor.
1850pub fn forward_embryo_graph(
1851    model: &Arc<EmbryoGraphModel>,
1852    kv_id: u64,
1853    hidden: &[f32],
1854    position: usize,
1855    logits: &mut Vec<f32>,
1856) -> bool {
1857    #[cfg(feature = "gpu")]
1858    if backend() == Backend::Wgpu {
1859        return crate::gpu_wgpu::forward_embryo_graph(model, kv_id, hidden, position, logits);
1860    }
1861    let _ = (model, kv_id, hidden, position, logits);
1862    false
1863}
1864
1865/// Speculative-verify tail for the batched graph: fold final-norm + lm_head
1866/// over every batch position and read all k logit rows back; the batch also
1867/// snapshots the GDN state per position for `gdn_spec_restore`.
1868#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1869pub enum BatchGraphOutcome {
1870    /// The graph declined before mutating persistent device state. Callers may
1871    /// safely use the existing per-position path.
1872    Declined,
1873    /// The complete batch committed and its readback succeeded.
1874    Completed,
1875    /// A batch that had admitted sealed O(1) state failed after admission.
1876    /// Falling back to CPU would mix two state machines, so the caller must
1877    /// abort and clear the sequence instead.
1878    Failed,
1879}
1880
1881pub struct SpecTail<'a> {
1882    pub lm: GraphW<'a>,
1883    pub lm_rows: usize,
1884    pub final_norm: &'a [f32],
1885    pub logits_out: &'a mut Vec<f32>,
1886}
1887
1888/// Batched prefill: k contiguous positions through the whole graph in one submit
1889/// (projections/FFN as GEMMs, attention/GDN looped over scratch). `h` is
1890/// [k·hidden] in/out; `positions` len k. wgpu only.
1891#[allow(clippy::too_many_arguments)]
1892pub fn forward_batch_graph(
1893    model: &Arc<CmfModel>,
1894    kv_id: u64,
1895    layers: &[GraphLayer],
1896    invf: &[f32],
1897    h: &mut [f32],
1898    nh: usize,
1899    nkv: usize,
1900    hd: usize,
1901    rd: usize,
1902    hidden: usize,
1903    inter: usize,
1904    positions: &[usize],
1905    cap: usize,
1906    gemma: bool,
1907    eps: f32,
1908    attn_scale: f32,
1909    k: usize,
1910    // Per-layer sealed O(1) device views. An empty slice means the ordinary
1911    // exact-KV path; otherwise it must have one entry per graph layer.
1912    o1: &[Option<Vec<crate::nystrom::O1DeviceView<'_>>>],
1913    o1_epoch: u64,
1914    spec: Option<SpecTail<'_>>,
1915    // Device-prefix mode (plain prefill only): Some = when the whole stack
1916    // does not fit the weight budget, run the leading layers that do — the
1917    // same prefix rule as the token graph — leave the boundary hidden in
1918    // `h` and report the count here; the caller runs the rest on the host.
1919    // None = all layers or a decline, as before.
1920    layers_run: Option<&mut usize>,
1921) -> BatchGraphOutcome {
1922    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)
1923}
1924
1925thread_local! {
1926    static MIMO_ATTN_SCRATCH: std::cell::Cell<bool> = const { std::cell::Cell::new(true) };
1927}
1928
1929pub(crate) fn mimo_attention_scratch_enabled() -> bool {
1930    MIMO_ATTN_SCRATCH.with(std::cell::Cell::get)
1931}
1932
1933/// Scoped diagnostic A/B switch; unlike process environment mutations it
1934/// cannot race a model's background workers. Restores on panic as well.
1935#[doc(hidden)]
1936pub fn mimo_attention_scratch_scope<R>(enabled: bool, f: impl FnOnce() -> R) -> R {
1937    struct Restore(bool);
1938    impl Drop for Restore {
1939        fn drop(&mut self) {
1940            MIMO_ATTN_SCRATCH.with(|v| v.set(self.0));
1941        }
1942    }
1943    let _restore = Restore(MIMO_ATTN_SCRATCH.with(|v| v.replace(enabled)));
1944    f()
1945}
1946
1947thread_local! {
1948    static MIMO_Q8_SHORT: std::cell::Cell<bool> = const { std::cell::Cell::new(true) };
1949}
1950
1951pub(crate) fn mimo_q8_short_enabled() -> bool {
1952    MIMO_Q8_SHORT.with(std::cell::Cell::get)
1953}
1954
1955/// Diagnostic A/B switch for row-exact short q8 graph kernels.
1956#[doc(hidden)]
1957pub fn mimo_q8_short_scope<R>(enabled: bool, f: impl FnOnce() -> R) -> R {
1958    struct Restore(bool);
1959    impl Drop for Restore {
1960        fn drop(&mut self) {
1961            MIMO_Q8_SHORT.with(|v| v.set(self.0));
1962        }
1963    }
1964    let _restore = Restore(MIMO_Q8_SHORT.with(|v| v.replace(enabled)));
1965    f()
1966}
1967
1968/// Batched graph over a span whose first absolute layer is `layer_base`.
1969#[allow(clippy::too_many_arguments)]
1970pub fn forward_batch_graph_at(
1971    model: &Arc<CmfModel>,
1972    kv_id: u64,
1973    layer_base: usize,
1974    layers: &[GraphLayer],
1975    invf: &[f32],
1976    h: &mut [f32],
1977    nh: usize,
1978    nkv: usize,
1979    hd: usize,
1980    rd: usize,
1981    hidden: usize,
1982    inter: usize,
1983    positions: &[usize],
1984    cap: usize,
1985    gemma: bool,
1986    eps: f32,
1987    attn_scale: f32,
1988    k: usize,
1989    // Per-layer sealed O(1) device views. An empty slice means the ordinary
1990    // exact-KV path; otherwise it must have one entry per graph layer.
1991    o1: &[Option<Vec<crate::nystrom::O1DeviceView<'_>>>],
1992    o1_epoch: u64,
1993    spec: Option<SpecTail<'_>>,
1994    // Device-prefix mode (plain prefill only): Some = when the whole stack
1995    // does not fit the weight budget, run the leading layers that do — the
1996    // same prefix rule as the token graph — leave the boundary hidden in
1997    // `h` and report the count here; the caller runs the rest on the host.
1998    // None = all layers or a decline, as before.
1999    layers_run: Option<&mut usize>,
2000) -> BatchGraphOutcome {
2001    match backend() {
2002        #[cfg(feature = "gpu")]
2003        Backend::Wgpu => crate::gpu_wgpu::forward_batch_graph_at(
2004            model, kv_id, layer_base, layers, invf, h, nh, nkv, hd, rd, hidden, inter, positions,
2005            cap, gemma, eps, attn_scale, k, o1, o1_epoch, spec, layers_run,
2006        ),
2007        #[allow(unreachable_patterns)]
2008        _ => {
2009            let _ = (o1, o1_epoch, spec, layers_run);
2010            BatchGraphOutcome::Declined
2011        }
2012    }
2013}
2014
2015/// After a partial speculative acceptance: restore every GDN layer's device
2016/// state to the snapshot after batch position `slot`. `base_pos` is the
2017/// absolute position of the first verify row and `expected_layers` makes the
2018/// restore all-or-nothing across the model's recurrent layers. wgpu only.
2019pub fn gdn_spec_restore(kv_id: u64, slot: usize, base_pos: usize, expected_layers: usize) -> bool {
2020    #[cfg(feature = "gpu")]
2021    if backend() == Backend::Wgpu {
2022        return crate::gpu_wgpu::gdn_spec_restore(kv_id, slot, base_pos, expected_layers);
2023    }
2024    #[allow(unreachable_code)]
2025    {
2026        let _ = (kv_id, slot, base_pos, expected_layers);
2027        false
2028    }
2029}
2030
2031/// Re-point one exact-attention device mirror after a speculative round has
2032/// discarded unaccepted rows. The rows beyond `stored` remain allocated and
2033/// are overwritten by the next append; only the logical cursor moves. This
2034/// is the wgpu twin of Metal's existing mirror cursor helper and keeps the
2035/// MTP graph's speculative/device cache coherent with its real anchor.
2036pub fn graph_kv_set_stored(kv_id: u64, layer: usize, stored: usize) -> bool {
2037    #[cfg(feature = "gpu")]
2038    if backend() == Backend::Wgpu {
2039        return crate::gpu_wgpu::kv_mirror_set_stored(kv_id, layer, stored);
2040    }
2041    #[cfg(target_os = "macos")]
2042    if backend() == Backend::Metal {
2043        crate::gpu_metal::kv_mirror_set_stored(kv_id, layer, stored);
2044        return true;
2045    }
2046    false
2047}
2048
2049/// Rows the wgpu token graph's exact-attention mirror holds for one layer
2050/// (None: no wgpu mirror). Metal keeps its owner cache current per token
2051/// and reports None here.
2052pub fn graph_kv_stored(_kv_id: u64, _layer: usize) -> Option<usize> {
2053    #[cfg(feature = "gpu")]
2054    if backend() == Backend::Wgpu {
2055        return crate::gpu_wgpu::kv_mirror_stored(_kv_id, _layer);
2056    }
2057    None
2058}
2059
2060/// Does the wgpu token graph hold a device-resident recurrent state for
2061/// this layer (one the host `linear_state` has not seen)?
2062pub fn graph_state_resident(_kv_id: u64, _layer: usize) -> bool {
2063    #[cfg(feature = "gpu")]
2064    if backend() == Backend::Wgpu {
2065        return crate::gpu_wgpu::graph_state_resident(_kv_id, _layer);
2066    }
2067    false
2068}
2069
2070/// Copy rows back from the wgpu token graph's K/V mirrors in one submit:
2071/// for each `(layer, from, to)` the K and V rows `[from..to)`, position-major
2072/// (`[(to − from) × nkv × hd]` each).
2073pub fn graph_kv_read_rows(
2074    _kv_id: u64,
2075    _reqs: &[(usize, usize, usize)],
2076    _nkv: usize,
2077    _hd: usize,
2078) -> Option<Vec<(Vec<f32>, Vec<f32>)>> {
2079    #[cfg(feature = "gpu")]
2080    if backend() == Backend::Wgpu {
2081        return crate::gpu_wgpu::kv_mirror_read_rows(_kv_id, _reqs, _nkv, _hd);
2082    }
2083    None
2084}
2085
2086/// Rows `[from, to)` of one wgpu exact-attention mirror in the host
2087/// cache's layout (V zero-padded to `hd`), whatever the mirror's geometry
2088/// (narrow V, a sliding layer's ring). The third value is the first
2089/// position actually read: a ring returns zeros below it. None: no wgpu
2090/// mirror holding those rows at (nkv, hd).
2091pub fn graph_kv_pull_host(
2092    _kv_id: u64,
2093    _layer: usize,
2094    _from: usize,
2095    _to: usize,
2096    _nkv: usize,
2097    _hd: usize,
2098) -> Option<(Vec<f32>, Vec<f32>, usize)> {
2099    #[cfg(feature = "gpu")]
2100    if backend() == Backend::Wgpu {
2101        return crate::gpu_wgpu::kv_mirror_pull_host(_kv_id, _layer, _from, _to, _nkv, _hd);
2102    }
2103    None
2104}
2105
2106/// Drop the wgpu token graph's device K/V mirror for a pipeline.
2107pub fn graph_kv_reset(_kv_id: u64) {
2108    #[cfg(feature = "gpu")]
2109    if backend() == Backend::Wgpu {
2110        crate::gpu_wgpu::kv_mirror_reset(_kv_id);
2111        crate::gpu_wgpu::embryo_graph_reset(_kv_id);
2112    }
2113}
2114
2115/// Ternary (q1t) BASE matvec on the GPU — fills `out` with the base dot; the
2116/// caller adds the sparse overlay on the CPU. Metal only for now (wgpu q1t not
2117/// yet written → CPU fallback).
2118pub fn q1t_matvec(
2119    model: &Arc<CmfModel>,
2120    idx: usize,
2121    xs: &[f32],
2122    rows: usize,
2123    cols: usize,
2124    out: &mut [f32],
2125) -> bool {
2126    match backend() {
2127        #[cfg(target_os = "macos")]
2128        Backend::Metal => {
2129            if metal_q1t_enabled() {
2130                crate::gpu_metal::q1t_matvec(model, idx, xs, rows, cols, out)
2131            } else {
2132                false
2133            }
2134        }
2135        #[cfg(feature = "gpu")]
2136        Backend::Wgpu => crate::gpu_wgpu::q1t_matvec(model, idx, xs, rows, cols, out),
2137        Backend::None => false,
2138    }
2139}
2140
2141/// q4_block matvec on the GPU — wgpu only (Metal drives q4_block through the
2142/// whole-token graph, not a standalone matvec).
2143#[allow(unused_variables)]
2144pub fn q4b_matvec(
2145    model: &Arc<CmfModel>,
2146    idx: usize,
2147    xs: &[f32],
2148    rows: usize,
2149    cols: usize,
2150    out: &mut [f32],
2151) -> bool {
2152    match backend() {
2153        #[cfg(target_os = "macos")]
2154        Backend::Metal => false,
2155        #[cfg(feature = "gpu")]
2156        Backend::Wgpu => crate::gpu_wgpu::q4b_matvec(model, idx, xs, rows, cols, out),
2157        Backend::None => false,
2158    }
2159}
2160
2161/// q1t batched GEMM (prefill) — base + overlay on-device (Metal simdgroup or
2162/// wgpu register-blocked).
2163pub fn q1t_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(target_os = "macos")]
2174        // Batched prefill and single-token decode are both enabled. On the
2175        // real 14.8B Q1T model prefill PPL was within 0.3% of CPU (7.942 vs
2176        // 7.966), and the alignment-safe decode kernel reached 3.52e-6 max_rel.
2177        Backend::Metal => crate::gpu_metal::q1t_matmat(model, idx, xs, b, rows, cols, out),
2178        #[cfg(feature = "gpu")]
2179        Backend::Wgpu => crate::gpu_wgpu::q1t_matmat(model, idx, xs, b, rows, cols, out),
2180        Backend::None => false,
2181    }
2182}
2183
2184/// Native Metal Q1T switch. Enabled by default after the byte-packed Q1T
2185/// fields were changed to alignment-safe loads; keep an explicit emergency
2186/// fallback for device/driver diagnostics.
2187#[cfg(target_os = "macos")]
2188pub(crate) fn metal_q1t_enabled() -> bool {
2189    std::env::var("CMF_METAL_Q1T")
2190        .map(|v| v != "0" && !v.eq_ignore_ascii_case("off"))
2191        .unwrap_or(true)
2192}
2193
2194/// Batched q1 GEMM (prefill). wgpu only — Metal has its own block path.
2195pub fn q1_matmat(
2196    model: &Arc<CmfModel>,
2197    idx: usize,
2198    xs: &[f32],
2199    b: usize,
2200    rows: usize,
2201    cols: usize,
2202    out: &mut [f32],
2203) -> bool {
2204    match backend() {
2205        #[cfg(feature = "gpu")]
2206        Backend::Wgpu => crate::gpu_wgpu::q1_matmat(model, idx, xs, b, rows, cols, out),
2207        #[allow(unused_variables)]
2208        _ => false,
2209    }
2210}
2211
2212/// Contention kill for the wide imagegen GEMM/FFN paths: one grossly
2213/// slow op under a work-proportional budget (fair-device ops are
2214/// ≤~100 ms even at 1024px) means another process owns the device —
2215/// verdicts are per-process, so CPU for the rest of this one.
2216static MM_KILL: AtomicBool = AtomicBool::new(false);
2217pub(crate) fn mm_killed() -> bool {
2218    MM_KILL.load(Ordering::Relaxed)
2219}
2220pub(crate) fn mm_kill() {
2221    MM_KILL.store(true, Ordering::Relaxed);
2222}
2223
2224/// Consecutive over-budget ops. ONE slow op is not contention: on a
2225/// 24 GB Mac running the 25.7 GB fl2va file the first ops after the
2226/// prompt encode page their weights in from the SSD and take seconds —
2227/// a field report (hololabs, HF discussion #2) had to neuter the kill
2228/// to keep the denoise on the GPU, and then measured 48 s/step where the
2229/// CPU fallback took >60. Contention is persistent; a page-in is not.
2230static MM_STRIKES: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
2231const MM_STRIKES_TO_KILL: u32 = 3;
2232/// Whether the kill is armed at all. A one-shot phase whose slowness is
2233/// expected and not contention — the video prompt encoder streaming
2234/// 12 GB off the SSD on a 24 GB Mac (HF discussion #4: users had to
2235/// gut `mm_kill` to keep the denoise loop on the GPU) — disarms it and
2236/// re-arms it when the phase is over; strikes taken meanwhile are
2237/// forgotten.
2238static MM_ARMED: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
2239
2240/// Disarm / re-arm the contention kill around a phase whose GEMMs are
2241/// slow for reasons that are not another process (see `MM_ARMED`).
2242pub fn mm_kill_arm(on: bool) {
2243    MM_ARMED.store(on, Ordering::Relaxed);
2244    if on {
2245        MM_STRIKES.store(0, Ordering::Relaxed);
2246    }
2247}
2248
2249/// The contention verdict for one wide op: `el` against its
2250/// work-proportional `budget`. `exempt` marks ops whose time is not
2251/// evidence — the cold probe, or a weight that was not resident before
2252/// the call and rode in with it. Kills after `MM_STRIKES_TO_KILL`
2253/// consecutive strikes; a within-budget op clears the count.
2254/// `CMF_MM_KILL=0` disables the kill entirely (the device is trusted).
2255pub(crate) fn mm_budget_check(
2256    what: &str,
2257    el: std::time::Duration,
2258    budget: std::time::Duration,
2259    exempt: bool,
2260) {
2261    if el <= budget {
2262        MM_STRIKES.store(0, Ordering::Relaxed);
2263        return;
2264    }
2265    if exempt || !MM_ARMED.load(Ordering::Relaxed) {
2266        return;
2267    }
2268    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
2269    let on = *ON.get_or_init(|| std::env::var("CMF_MM_KILL").as_deref() != Ok("0"));
2270    let n = MM_STRIKES.fetch_add(1, Ordering::Relaxed) + 1;
2271    if !on {
2272        tracing::info!(
2273            "gpu {what} took {el:?} (budget {budget:?}) — over budget, CMF_MM_KILL=0 keeps the device"
2274        );
2275        return;
2276    }
2277    if n >= MM_STRIKES_TO_KILL {
2278        tracing::warn!(
2279            "gpu {what} took {el:?} (budget {budget:?}), {n} in a row — \
2280             device contended, CPU for the rest of the process (CMF_MM_KILL=0 to override)"
2281        );
2282        mm_kill();
2283    } else {
2284        tracing::info!(
2285            "gpu {what} took {el:?} (budget {budget:?}) — strike {n} of {MM_STRIKES_TO_KILL}"
2286        );
2287    }
2288}
2289
2290/// Fused DiT SwiGLU FFN on the device: g=X·W1ᵀ, u=X·W3ᵀ, silu(g)·u,
2291/// Causal chunk attention on the device: `b` queries against `s0 + b`
2292/// cached keys. wgpu only — Metal's chunk graph keeps attention inside
2293/// the resident block and never calls out.
2294#[allow(unused_variables, clippy::too_many_arguments)]
2295pub fn chunk_attend(
2296    q: &[f32],
2297    k: &[&[f32]],
2298    v: &[&[f32]],
2299    b: usize,
2300    s0: usize,
2301    nh: usize,
2302    nkv: usize,
2303    hd: usize,
2304    scale: f32,
2305    out: &mut [f32],
2306) -> bool {
2307    match backend() {
2308        #[cfg(feature = "gpu")]
2309        Backend::Wgpu => crate::gpu_wgpu::chunk_attend(q, k, v, b, s0, nh, nkv, hd, scale, out),
2310        #[allow(unreachable_patterns)]
2311        _ => false,
2312    }
2313}
2314
2315/// `chunk_attend` under a sliding window (`window` keys, the query's own
2316/// included). wgpu only; see `chunk_attend_windowed`.
2317#[allow(unused_variables, clippy::too_many_arguments)]
2318pub fn chunk_attend_win(
2319    q: &[f32],
2320    k: &[&[f32]],
2321    v: &[&[f32]],
2322    b: usize,
2323    s0: usize,
2324    nh: usize,
2325    nkv: usize,
2326    hd: usize,
2327    scale: f32,
2328    window: usize,
2329    out: &mut [f32],
2330) -> bool {
2331    match backend() {
2332        #[cfg(feature = "gpu")]
2333        Backend::Wgpu => {
2334            crate::gpu_wgpu::chunk_attend_win(q, k, v, b, s0, nh, nkv, hd, scale, window, out)
2335        }
2336        #[allow(unreachable_patterns)]
2337        _ => false,
2338    }
2339}
2340
2341/// Several projections of one host activation, read back in one trip:
2342/// `out` = [P0 | P1 | …], each `idxs` entry (tensor, rows). wgpu only.
2343#[allow(unused_variables)]
2344pub fn gemm_many_keep(
2345    model: &Arc<CmfModel>,
2346    idxs: &[(usize, usize)],
2347    xs: &[f32],
2348    b: usize,
2349    cols: usize,
2350    out: &mut [f32],
2351) -> bool {
2352    match backend() {
2353        #[cfg(feature = "gpu")]
2354        Backend::Wgpu => crate::gpu_wgpu::gemm_many_keep(model, idxs, xs, b, cols, out),
2355        #[allow(unreachable_patterns)]
2356        _ => false,
2357    }
2358}
2359
2360thread_local! {
2361    /// Set by the pipeline around a layer of Spark-X2.5's batched prefill:
2362    /// its int8 matrix-unit GEMMs may take the `zi_mm` tile, and the
2363    /// projections that read one activation share its upload (wgpu).
2364    static PREFILL_FAST_GEMM: Cell<bool> = const { Cell::new(false) };
2365}
2366
2367/// Restores the previous flag on drop.
2368pub struct PrefillFastGemmGuard(bool);
2369
2370impl Drop for PrefillFastGemmGuard {
2371    fn drop(&mut self) {
2372        PREFILL_FAST_GEMM.with(|c| c.set(self.0));
2373    }
2374}
2375
2376/// Let this thread's int8 matrix-unit GEMMs take the `zi_mm` tile for the
2377/// guard's lifetime (bit-identical to the 64x64 kernel; see
2378/// `gpu_wgpu::zi_gemm_route`).
2379pub fn enter_prefill_fast_gemm() -> PrefillFastGemmGuard {
2380    PrefillFastGemmGuard(PREFILL_FAST_GEMM.with(|c| c.replace(true)))
2381}
2382
2383/// Whether `enter_prefill_fast_gemm` is in force on this thread.
2384pub fn prefill_fast_gemm() -> bool {
2385    PREFILL_FAST_GEMM.with(|c| c.get())
2386}
2387
2388/// The largest finite f16. An operand the matrix-unit attention casts to
2389/// f16 must stay within it, or the cast overflows to inf.
2390pub const F16_MAX: f32 = 65504.0;
2391
2392/// max |x| over `xs`, or +inf when an entry is inf or NaN: the value the
2393/// f16 operand guards compare with `F16_MAX`. Sixteen independent lanes
2394/// and a flag, so it vectorizes (a single `fold` chain does not).
2395pub fn abs_max_or_inf(xs: &[f32]) -> f32 {
2396    let mut acc = [0f32; 16];
2397    let mut bad = [false; 16];
2398    let mut it = xs.chunks_exact(16);
2399    for c in &mut it {
2400        for ((a, f), &v) in acc.iter_mut().zip(bad.iter_mut()).zip(c) {
2401            let v = v.abs();
2402            // NaN compares false both ways: flagged, never taken.
2403            *f |= !(v < f32::INFINITY);
2404            *a = if v > *a { v } else { *a };
2405        }
2406    }
2407    let mut m = 0f32;
2408    for &v in it.remainder() {
2409        let v = v.abs();
2410        if !(v < f32::INFINITY) {
2411            return f32::INFINITY;
2412        }
2413        m = m.max(v);
2414    }
2415    if bad.iter().any(|&f| f) {
2416        return f32::INFINITY;
2417    }
2418    acc.iter().fold(m, |m, &a| m.max(a))
2419}
2420
2421/// The decode graph's device K/V mirror a prefill chunk of one layer also
2422/// writes to (wgpu): `kv_id` is the pipeline's graph id, `layer` the
2423/// virtual layer index, `limit` the cache's `max_seq_len` — the same key
2424/// and capacity ceiling the token graph uses, so the mirror the prefill
2425/// fills is the one the first decode token reads.
2426#[derive(Clone, Copy, Debug, PartialEq, Eq)]
2427pub struct PrefillMirror {
2428    pub kv_id: u64,
2429    pub layer: usize,
2430    pub limit: usize,
2431}
2432
2433thread_local! {
2434    /// Set by the pipeline around one layer's batched attention (see
2435    /// `enter_prefill_mirror`); every other caller of the batched
2436    /// attention (MTP heads, tests) sees None and keeps the host upload.
2437    static PREFILL_MIRROR: Cell<Option<PrefillMirror>> = const { Cell::new(None) };
2438}
2439
2440/// Restores the previous target on drop.
2441pub struct PrefillMirrorGuard(Option<PrefillMirror>);
2442
2443impl Drop for PrefillMirrorGuard {
2444    fn drop(&mut self) {
2445        PREFILL_MIRROR.with(|c| c.set(self.0));
2446    }
2447}
2448
2449/// Name the device mirror the current layer's prefill chunk appends to
2450/// and attends against, for the guard's lifetime.
2451pub fn enter_prefill_mirror(t: PrefillMirror) -> PrefillMirrorGuard {
2452    PrefillMirrorGuard(PREFILL_MIRROR.with(|c| c.replace(Some(t))))
2453}
2454
2455/// The mirror `enter_prefill_mirror` named on this thread, if any.
2456pub fn prefill_mirror() -> Option<PrefillMirror> {
2457    PREFILL_MIRROR.with(|c| c.get())
2458}
2459
2460/// `chunk_attend_win` against the decode graph's device mirror instead of
2461/// an upload of the whole prefix: the host rows the mirror lacks (the
2462/// chunk's own `b`, or more after a host-only chunk) are appended first,
2463/// then the attention binds the mirror rows in place. `s0` is the chunk's
2464/// first position and `base` the position of host row 0
2465/// (`LayerKvCache::base`: past 0 on a trimmed sliding tail). `ring` is the
2466/// layer's window as the mirror geometry (Some on every sliding layer),
2467/// `window` the window the softmax masks with (0 while it masks nothing).
2468/// `operand_max` = max |x| over `q` and every K/V row the attend reads
2469/// (`abs_max_or_inf`; +inf when unknown): past `F16_MAX` the f32 kernels
2470/// run instead of the matrix units. wgpu only; false = nothing attended
2471/// (the caller uploads as before).
2472#[allow(unused_variables, clippy::too_many_arguments)]
2473pub fn chunk_attend_mirror(
2474    t: PrefillMirror,
2475    cpu_k: &[Vec<f32>],
2476    cpu_v: &[Vec<f32>],
2477    base: usize,
2478    q: &[f32],
2479    b: usize,
2480    s0: usize,
2481    nh: usize,
2482    nkv: usize,
2483    hd: usize,
2484    scale: f32,
2485    ring: Option<usize>,
2486    window: usize,
2487    operand_max: f32,
2488    out: &mut [f32],
2489) -> bool {
2490    match backend() {
2491        #[cfg(feature = "gpu")]
2492        // QKᵀ and P·V on the matrix units where the device has them (f16
2493        // operands, f32 accumulators — not bit-identical to the f32
2494        // kernels: Spark-X2.5 4B q8_2f wiki ppl 10.677 -> 10.672);
2495        // `CMF_PREFILL_ATTN_COOP=0` keeps the f32 ones (A/B).
2496        Backend::Wgpu => crate::gpu_wgpu::chunk_attend_mirror(
2497            t.kv_id,
2498            t.layer,
2499            t.limit,
2500            cpu_k,
2501            cpu_v,
2502            base,
2503            q,
2504            b,
2505            s0,
2506            nh,
2507            nkv,
2508            hd,
2509            scale,
2510            ring,
2511            window,
2512            std::env::var("CMF_PREFILL_ATTN_COOP").as_deref() != Ok("0"),
2513            operand_max,
2514            out,
2515        ),
2516        #[allow(unreachable_patterns)]
2517        _ => false,
2518    }
2519}
2520
2521/// `chunk_attend_mirror`, then the head gate (`gains`, b × nh, computed on
2522/// the host) and the O projection `wo` on the card: the attention output
2523/// never comes home. `out` = [b][hidden]. The card scans the O operand
2524/// itself, which is the host path's product bit for bit only while max
2525/// |gated output| <= 1000: the caller vouches for that (Spark-X2.5's
2526/// prefill, max |v| <= 990). wgpu only; false = nothing written to `out`
2527/// (the caller attends and projects as before).
2528#[allow(unused_variables, clippy::too_many_arguments)]
2529pub fn chunk_attend_mirror_wo(
2530    t: PrefillMirror,
2531    cpu_k: &[Vec<f32>],
2532    cpu_v: &[Vec<f32>],
2533    base: usize,
2534    q: &[f32],
2535    b: usize,
2536    s0: usize,
2537    nh: usize,
2538    nkv: usize,
2539    hd: usize,
2540    scale: f32,
2541    ring: Option<usize>,
2542    window: usize,
2543    operand_max: f32,
2544    gains: Option<&[f32]>,
2545    wo: (&Arc<CmfModel>, usize),
2546    hidden: usize,
2547    out: &mut [f32],
2548) -> bool {
2549    match backend() {
2550        #[cfg(feature = "gpu")]
2551        Backend::Wgpu => crate::gpu_wgpu::chunk_attend_mirror_wo(
2552            t.kv_id,
2553            t.layer,
2554            t.limit,
2555            cpu_k,
2556            cpu_v,
2557            base,
2558            q,
2559            b,
2560            s0,
2561            nh,
2562            nkv,
2563            hd,
2564            scale,
2565            ring,
2566            window,
2567            std::env::var("CMF_PREFILL_ATTN_COOP").as_deref() != Ok("0"),
2568            operand_max,
2569            gains,
2570            wo,
2571            hidden,
2572            out,
2573        ),
2574        #[allow(unreachable_patterns)]
2575        _ => false,
2576    }
2577}
2578
2579/// Whether the active backend's chunk attend can apply a sliding window.
2580pub fn chunk_attend_windowed() -> bool {
2581    match backend() {
2582        #[cfg(feature = "gpu")]
2583        Backend::Wgpu => true,
2584        #[allow(unreachable_patterns)]
2585        _ => false,
2586    }
2587}
2588
2589/// Fused QKV projection: one upload of the normed chunk, three GEMMs,
2590/// one readback of Q|K|V back to back. Metal has no twin yet — its
2591/// chunk graph keeps the whole layer resident and never surfaces QKV.
2592#[allow(unused_variables, clippy::too_many_arguments)]
2593pub fn q4t_qkv(
2594    model: &Arc<CmfModel>,
2595    wq: usize,
2596    wk: usize,
2597    wv: usize,
2598    xs: &[f32],
2599    b: usize,
2600    cols: usize,
2601    rq: usize,
2602    rk: usize,
2603    rv: usize,
2604    out: &mut [f32],
2605) -> bool {
2606    match backend() {
2607        #[cfg(feature = "gpu")]
2608        Backend::Wgpu => crate::gpu_wgpu::q4t_qkv(model, wq, wk, wv, xs, b, cols, rq, rk, rv, out),
2609        #[allow(unreachable_patterns)]
2610        _ => false,
2611    }
2612}
2613
2614/// y=·W2ᵀ — one command buffer, only X and Y cross the CPU boundary.
2615#[allow(unused_variables, clippy::too_many_arguments)]
2616/// SwiGLU FFN with a row-packed [gate|up] fc1 (MiniMax-H3's DiT), run
2617/// end to end on the device. wgpu only: Metal keeps the host loop until
2618/// its own packed kernel exists.
2619#[allow(clippy::too_many_arguments, unused_variables)]
2620pub fn q4tp_ffn_packed(
2621    model: &Arc<CmfModel>,
2622    w1: usize,
2623    w2: usize,
2624    xs: &[f32],
2625    b: usize,
2626    hidden: usize,
2627    inter: usize,
2628    bias: Option<&[f32]>,
2629    out: &mut [f32],
2630) -> bool {
2631    match backend() {
2632        #[cfg(feature = "gpu")]
2633        Backend::Wgpu => {
2634            crate::gpu_wgpu::ffn_packed(model, w1, w2, xs, b, hidden, inter, bias, out)
2635        }
2636        #[allow(unreachable_patterns)]
2637        _ => false,
2638    }
2639}
2640
2641pub fn q4tp_ffn(
2642    model: &Arc<CmfModel>,
2643    w1: usize,
2644    w3: usize,
2645    w2: usize,
2646    xs: &[f32],
2647    b: usize,
2648    hidden: usize,
2649    inter: usize,
2650    out: &mut [f32],
2651) -> bool {
2652    match backend() {
2653        #[cfg(target_os = "macos")]
2654        Backend::Metal => crate::gpu_metal::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2655        #[cfg(feature = "gpu")]
2656        Backend::Wgpu => crate::gpu_wgpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2657        #[allow(unreachable_patterns)]
2658        _ => false,
2659    }
2660}
2661
2662/// `q4tp_ffn` (`q4tp` = true) / `q4t_ffn` with the activation named:
2663/// SiLU takes the plain entry points (every backend), the exact GELU of
2664/// Spark-X2.5 has a wgpu arm only — other backends decline it.
2665#[allow(clippy::too_many_arguments, unused_variables)]
2666pub fn q4_ffn_act(
2667    model: &Arc<CmfModel>,
2668    w1: usize,
2669    w3: usize,
2670    w2: usize,
2671    xs: &[f32],
2672    b: usize,
2673    hidden: usize,
2674    inter: usize,
2675    q4tp: bool,
2676    act: GraphAct,
2677    out: &mut [f32],
2678) -> bool {
2679    match act {
2680        GraphAct::Silu if q4tp => q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2681        GraphAct::Silu => q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2682        GraphAct::GeluErf => match backend() {
2683            #[cfg(feature = "gpu")]
2684            Backend::Wgpu => crate::gpu_wgpu::q4_ffn_act(
2685                model,
2686                w1,
2687                w3,
2688                w2,
2689                xs,
2690                b,
2691                hidden,
2692                inter,
2693                q4tp,
2694                act.code(),
2695                out,
2696            ),
2697            #[allow(unreachable_patterns)]
2698            _ => false,
2699        },
2700    }
2701}
2702
2703/// The LLM prefill's dense FFN for separate gate/up/down weights in any
2704/// codec with a device GEMM (int8, q4tp, or a mix of them): the panels
2705/// stay on the card between the projections. wgpu only; other backends
2706/// decline and the caller keeps its per-GEMM path.
2707#[allow(clippy::too_many_arguments, unused_variables)]
2708pub fn ffn_act_keep(
2709    model: &Arc<CmfModel>,
2710    w1: usize,
2711    w3: usize,
2712    w2: usize,
2713    xs: &[f32],
2714    b: usize,
2715    hidden: usize,
2716    inter: usize,
2717    act: GraphAct,
2718    out: &mut [f32],
2719) -> bool {
2720    match backend() {
2721        #[cfg(feature = "gpu")]
2722        Backend::Wgpu => crate::gpu_wgpu::ffn_act_keep(
2723            model,
2724            w1,
2725            w3,
2726            w2,
2727            xs,
2728            b,
2729            hidden,
2730            inter,
2731            act.code(),
2732            out,
2733        ),
2734        #[allow(unreachable_patterns)]
2735        _ => false,
2736    }
2737}
2738
2739/// Qwen Image's exact two-projection tanh-GELU FFN.  The WGPU arm keeps the
2740/// intermediate on the device; other backends decline so the caller retains
2741/// its bounded CPU path.  `bias_in` is applied before GELU and `bias_out`
2742/// after the second projection, matching the official transformer.
2743#[allow(clippy::too_many_arguments, unused_variables)]
2744pub fn q4tp_gelu_ffn(
2745    model: &Arc<CmfModel>,
2746    w_in: usize,
2747    w_out: usize,
2748    xs: &[f32],
2749    b: usize,
2750    hidden: usize,
2751    inter: usize,
2752    bias_in: &[f32],
2753    bias_out: &[f32],
2754    out: &mut [f32],
2755) -> bool {
2756    match backend() {
2757        #[cfg(feature = "gpu")]
2758        Backend::Wgpu => crate::gpu_wgpu::q4tp_gelu_ffn(
2759            model, w_in, w_out, xs, b, hidden, inter, bias_in, bias_out, out,
2760        ),
2761        #[allow(unreachable_patterns)]
2762        _ => false,
2763    }
2764}
2765
2766pub fn q4t_ffn(
2767    model: &Arc<CmfModel>,
2768    w1: usize,
2769    w3: usize,
2770    w2: usize,
2771    xs: &[f32],
2772    b: usize,
2773    hidden: usize,
2774    inter: usize,
2775    out: &mut [f32],
2776) -> bool {
2777    match backend() {
2778        #[cfg(target_os = "macos")]
2779        Backend::Metal => crate::gpu_metal::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2780        #[cfg(feature = "gpu")]
2781        Backend::Wgpu => crate::gpu_wgpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, out),
2782        #[allow(unreachable_patterns)]
2783        _ => false,
2784    }
2785}
2786
2787/// One whole modulated DiT block for `dit_block`: geometry, norm
2788/// weights, AdaLN scale/gate vectors (gates pre-tanh'd), a per-token
2789/// f32 RoPE cos/sin table, and the directory indices of the seven
2790/// q4t projections. `x` is in-out `[n, hidden]`.
2791pub struct DitBlockArgs<'a> {
2792    pub n: usize,
2793    pub hidden: usize,
2794    pub inter: usize,
2795    pub nh: usize,
2796    pub nkv: usize,
2797    pub hd: usize,
2798    pub eps: f32,
2799    pub rope_cos: &'a [f32],
2800    pub rope_sin: &'a [f32],
2801    pub norm1: &'a [f32],
2802    pub norm2: &'a [f32],
2803    pub ffn_norm1: &'a [f32],
2804    pub ffn_norm2: &'a [f32],
2805    pub norm_q: &'a [f32],
2806    pub norm_k: &'a [f32],
2807    pub s_msa: &'a [f32],
2808    pub gate_msa: &'a [f32],
2809    pub s_mlp: &'a [f32],
2810    pub gate_mlp: &'a [f32],
2811    pub wq: usize,
2812    pub wk: usize,
2813    pub wv: usize,
2814    pub wo: usize,
2815    pub w1: usize,
2816    pub w3: usize,
2817    pub w2: usize,
2818    /// The projections' layout: q4tp (ladder scales) vs plain q4_tiled.
2819    /// The recommended Lumina file is q4tp, and a backend that only
2820    /// knows q4t must decline rather than decode with the wrong reader.
2821    pub q4tp: bool,
2822    /// The hidden state is already on the device from the previous block,
2823    /// so `x` need not be uploaded.
2824    pub resident_in: bool,
2825    /// Leave the result on the device instead of reading it back. The DiT
2826    /// loop does not touch `x` between blocks, so 27 of every 28 readbacks
2827    /// were moving 19 MB across PCIe and stalling on it for nothing.
2828    pub resident_out: bool,
2829}
2830
2831/// Can the selected backend keep the DiT's hidden state on the device
2832/// between blocks? Only the wgpu whole-block path; the Metal entry takes
2833/// and returns host memory every call.
2834pub fn dit_chain_supported() -> bool {
2835    #[cfg(feature = "gpu")]
2836    {
2837        return matches!(backend(), Backend::Wgpu) && fused_dit_block_available();
2838    }
2839    #[allow(unreachable_code)]
2840    false
2841}
2842
2843/// Pull the resident hidden state back to the host. For the caller that
2844/// chained blocks and then hit one the device declined.
2845pub fn dit_state_fetch(_x: &mut [f32]) -> bool {
2846    #[cfg(feature = "gpu")]
2847    {
2848        if matches!(backend(), Backend::Wgpu) {
2849            return crate::gpu_wgpu::dit_state_fetch(_x);
2850        }
2851    }
2852    false
2853}
2854
2855/// One whole modulated DiT block on the device — norms, qkv, RoPE,
2856/// attention, residuals and the SwiGLU FFN in a single command
2857/// buffer; only `x` crosses the CPU boundary (in and out).
2858#[allow(unused_variables)]
2859/// The DiT's three projections in one submission (wgpu only; the
2860/// Metal path fuses the whole block instead). False = the caller keeps
2861/// its three separate calls.
2862#[allow(unused_variables, clippy::too_many_arguments)]
2863pub fn dit_qkv(
2864    model: &Arc<CmfModel>,
2865    wq: usize,
2866    wk: usize,
2867    wv: usize,
2868    xs: &[f32],
2869    b: usize,
2870    hidden: usize,
2871    qrows: usize,
2872    kvrows: usize,
2873    q_out: &mut [f32],
2874    k_out: &mut [f32],
2875    v_out: &mut [f32],
2876) -> bool {
2877    match backend() {
2878        #[cfg(feature = "gpu")]
2879        Backend::Wgpu => crate::gpu_wgpu::q4tp_qkv(
2880            model, wq, wk, wv, xs, b, hidden, qrows, kvrows, q_out, k_out, v_out,
2881        ),
2882        #[allow(unreachable_patterns)]
2883        _ => false,
2884    }
2885}
2886
2887/// The Qwen Image double-stream attention half.  The WGPU implementation
2888/// keeps the six Q/K/V projections, the stream join, qk-norm/RoPE, joint
2889/// attention, and both output projections on the device; the caller only
2890/// supplies the two normalized streams and receives the two projected
2891/// streams.  A backend or codec that cannot satisfy the full contract
2892/// returns `false` before changing either output, so the native host path
2893/// remains the portable fallback.
2894pub struct QwenImageAttentionArgs<'a> {
2895    pub image: &'a [f32],
2896    pub text: &'a [f32],
2897    pub image_tokens: usize,
2898    pub text_tokens: usize,
2899    pub heads: usize,
2900    pub head_dim: usize,
2901    pub image_q: usize,
2902    pub image_k: usize,
2903    pub image_v: usize,
2904    pub text_q: usize,
2905    pub text_k: usize,
2906    pub text_v: usize,
2907    pub image_out: usize,
2908    pub text_out: usize,
2909    pub image_q_norm: &'a [f32],
2910    pub image_k_norm: &'a [f32],
2911    pub text_q_norm: &'a [f32],
2912    pub text_k_norm: &'a [f32],
2913    pub image_cos: &'a [f32],
2914    pub image_sin: &'a [f32],
2915    pub text_cos: &'a [f32],
2916    pub text_sin: &'a [f32],
2917    pub image_q_bias: &'a [f32],
2918    pub image_k_bias: &'a [f32],
2919    pub image_v_bias: &'a [f32],
2920    pub text_q_bias: &'a [f32],
2921    pub text_k_bias: &'a [f32],
2922    pub text_v_bias: &'a [f32],
2923    pub image_out_bias: &'a [f32],
2924    pub text_out_bias: &'a [f32],
2925    pub image_proj: &'a mut [f32],
2926    pub text_proj: &'a mut [f32],
2927}
2928
2929/// The per-layer controls and Q4TP directory indices used by the native
2930/// Qwen block.  Keeping this descriptor separate from the stream buffers
2931/// lets a whole transformer forward reuse one explicit device state without
2932/// a global scratch slot or a hidden context label.
2933#[allow(clippy::too_many_fields)]
2934pub struct QwenImageChainBlock<'a> {
2935    pub image_mod: &'a [f32],
2936    pub text_mod: &'a [f32],
2937    pub image_q: usize,
2938    pub image_k: usize,
2939    pub image_v: usize,
2940    pub text_q: usize,
2941    pub text_k: usize,
2942    pub text_v: usize,
2943    pub image_out: usize,
2944    pub text_out: usize,
2945    pub image_q_norm: &'a [f32],
2946    pub image_k_norm: &'a [f32],
2947    pub text_q_norm: &'a [f32],
2948    pub text_k_norm: &'a [f32],
2949    pub image_q_bias: &'a [f32],
2950    pub image_k_bias: &'a [f32],
2951    pub image_v_bias: &'a [f32],
2952    pub text_q_bias: &'a [f32],
2953    pub text_k_bias: &'a [f32],
2954    pub text_v_bias: &'a [f32],
2955    pub image_out_bias: &'a [f32],
2956    pub text_out_bias: &'a [f32],
2957    pub image_attn_gate: &'a [f32],
2958    pub text_attn_gate: &'a [f32],
2959    pub image_mlp_in: usize,
2960    pub image_mlp_out: usize,
2961    pub text_mlp_in: usize,
2962    pub text_mlp_out: usize,
2963    pub image_mlp_in_bias: &'a [f32],
2964    pub image_mlp_out_bias: &'a [f32],
2965    pub text_mlp_in_bias: &'a [f32],
2966    pub text_mlp_out_bias: &'a [f32],
2967}
2968
2969/// Complete Qwen Image transformer block contract. The first norm/mod
2970/// panels are supplied by the native caller; the WGPU arm keeps both streams
2971/// resident through QKV, QK/RoPE, joint attention, output projections, both
2972/// gated residuals, and the exact tanh-GELU MLPs. A backend that cannot
2973/// satisfy the whole graph returns `false` without changing either output.
2974#[allow(clippy::too_many_fields)]
2975pub struct QwenImageBlockArgs<'a> {
2976    /// Raw stream state is read for the first gated residual and overwritten
2977    /// with the block's final state after the one readback.
2978    pub image: &'a mut [f32],
2979    pub text: &'a mut [f32],
2980    pub image_norm: &'a [f32],
2981    pub text_norm: &'a [f32],
2982    pub image_tokens: usize,
2983    pub text_tokens: usize,
2984    pub heads: usize,
2985    pub head_dim: usize,
2986    pub image_cos: &'a [f32],
2987    pub image_sin: &'a [f32],
2988    pub text_cos: &'a [f32],
2989    pub text_sin: &'a [f32],
2990    pub image_q: usize,
2991    pub image_k: usize,
2992    pub image_v: usize,
2993    pub text_q: usize,
2994    pub text_k: usize,
2995    pub text_v: usize,
2996    pub image_out: usize,
2997    pub text_out: usize,
2998    pub image_q_norm: &'a [f32],
2999    pub image_k_norm: &'a [f32],
3000    pub text_q_norm: &'a [f32],
3001    pub text_k_norm: &'a [f32],
3002    pub image_q_bias: &'a [f32],
3003    pub image_k_bias: &'a [f32],
3004    pub image_v_bias: &'a [f32],
3005    pub text_q_bias: &'a [f32],
3006    pub text_k_bias: &'a [f32],
3007    pub text_v_bias: &'a [f32],
3008    pub image_out_bias: &'a [f32],
3009    pub text_out_bias: &'a [f32],
3010    pub image_attn_gate: &'a [f32],
3011    pub text_attn_gate: &'a [f32],
3012    pub image_mlp_in: usize,
3013    pub image_mlp_out: usize,
3014    pub text_mlp_in: usize,
3015    pub text_mlp_out: usize,
3016    pub image_mlp_in_bias: &'a [f32],
3017    pub image_mlp_out_bias: &'a [f32],
3018    pub text_mlp_in_bias: &'a [f32],
3019    pub text_mlp_out_bias: &'a [f32],
3020    pub image_mlp_mod: &'a [f32],
3021    pub text_mlp_mod: &'a [f32],
3022    pub image_mlp_gate: &'a [f32],
3023    pub text_mlp_gate: &'a [f32],
3024}
3025
3026/// Explicit whole-forward Qwen state contract.  The WGPU backend uploads the
3027/// two initial streams once, encodes a bounded number of complete blocks per
3028/// submission, and reads the final state once.  `blocks` is immutable for the
3029/// call, while the two stream slices receive only the final readback.
3030pub struct QwenImageChainArgs<'a> {
3031    pub image: &'a mut [f32],
3032    pub text: &'a mut [f32],
3033    pub image_tokens: usize,
3034    pub text_tokens: usize,
3035    pub heads: usize,
3036    pub head_dim: usize,
3037    pub image_cos: &'a [f32],
3038    pub image_sin: &'a [f32],
3039    pub text_cos: &'a [f32],
3040    pub text_sin: &'a [f32],
3041    pub blocks: &'a [QwenImageChainBlock<'a>],
3042}
3043
3044#[allow(unused_variables)]
3045pub fn qwen_image_attention(
3046    model: &Arc<CmfModel>,
3047    args: &mut QwenImageAttentionArgs<'_>,
3048) -> bool {
3049    match backend() {
3050        #[cfg(feature = "gpu")]
3051        Backend::Wgpu => crate::gpu_wgpu::qwen_image_attention(model, args),
3052        #[allow(unreachable_patterns)]
3053        _ => false,
3054    }
3055}
3056
3057#[allow(unused_variables)]
3058pub fn qwen_image_block(model: &Arc<CmfModel>, args: &mut QwenImageBlockArgs<'_>) -> bool {
3059    match backend() {
3060        #[cfg(feature = "gpu")]
3061        Backend::Wgpu => crate::gpu_wgpu::qwen_image_block(model, args),
3062        #[allow(unreachable_patterns)]
3063        _ => false,
3064    }
3065}
3066
3067/// Keep all Qwen transformer blocks on the selected WGPU device, with only
3068/// bounded chunk submissions and one final readback.  Other backends decline
3069/// so the native caller can use its exact portable block loop.
3070#[allow(unused_variables)]
3071pub fn qwen_image_chain(model: &Arc<CmfModel>, args: &mut QwenImageChainArgs<'_>) -> bool {
3072    match backend() {
3073        #[cfg(feature = "gpu")]
3074        Backend::Wgpu => crate::gpu_wgpu::qwen_image_chain(model, args),
3075        #[allow(unreachable_patterns)]
3076        _ => false,
3077    }
3078}
3079
3080/// The Qwen Image second sub-block on WGPU: affine-free LayerNorm,
3081/// shift/scale modulation, Q4TP input projection, exact tanh-GELU, output
3082/// projection, bias and gated residual.  `data` is updated in place after a
3083/// single final readback.  Backends/codecs that cannot keep this chain on the
3084/// device return `false` before changing `data`, leaving the caller's
3085/// portable per-op path intact.
3086#[allow(unused_variables, clippy::too_many_arguments)]
3087pub fn qwen_image_mlp_inplace(
3088    model: &Arc<CmfModel>,
3089    w_in: usize,
3090    w_out: usize,
3091    data: &mut [f32],
3092    batch: usize,
3093    hidden: usize,
3094    inter: usize,
3095    bias_in: &[f32],
3096    bias_out: &[f32],
3097    modulation: &[f32],
3098    gate: &[f32],
3099) -> bool {
3100    match backend() {
3101        #[cfg(feature = "gpu")]
3102        Backend::Wgpu => crate::gpu_wgpu::qwen_image_mlp_inplace(
3103            model,
3104            w_in,
3105            w_out,
3106            data,
3107            batch,
3108            hidden,
3109            inter,
3110            bias_in,
3111            bias_out,
3112            modulation,
3113            gate,
3114        ),
3115        #[allow(unreachable_patterns)]
3116        _ => false,
3117    }
3118}
3119
3120/// Is a FUSED whole-block device path on offer? The batched-CFG shape
3121/// (two sequences in one tall batch) and the fused block (one sequence,
3122/// one command buffer) are alternatives, and the caller picks.
3123pub fn fused_dit_block_available() -> bool {
3124    #[cfg(target_os = "macos")]
3125    {
3126        matches!(backend(), Backend::Metal) && fused_block_trusted()
3127    }
3128    #[cfg(not(target_os = "macos"))]
3129    {
3130        false
3131    }
3132}
3133
3134pub fn dit_block(model: &Arc<CmfModel>, a: &DitBlockArgs, x: &mut [f32]) -> bool {
3135    dit_block_seg(model, a, &[a.n], x)
3136}
3137
3138/// The same block over a CONCATENATION of independent sequences:
3139/// attention per segment, everything position-wise batched. wgpu only —
3140/// the Metal path takes the single-sequence entry above.
3141pub fn dit_block_seg(
3142    model: &Arc<CmfModel>,
3143    a: &DitBlockArgs,
3144    segs: &[usize],
3145    x: &mut [f32],
3146) -> bool {
3147    match backend() {
3148        #[cfg(target_os = "macos")]
3149        Backend::Metal if segs.len() <= 1 => crate::gpu_metal::dit_block(model, a, x),
3150        // The wgpu whole-block path. What it buys is host round trips —
3151        // six a block become one — so it defaults ON where those cost
3152        // real time (a discrete card across PCIe) and OFF on unified
3153        // memory, where the per-op path shares the same pages and the
3154        // fusion measured slightly slower on an M4. `CMF_DIT_FUSED=1`
3155        // forces it anywhere, `=0` forbids it.
3156        #[cfg(feature = "gpu")]
3157        Backend::Wgpu
3158            if match std::env::var("CMF_DIT_FUSED").ok().as_deref() {
3159                Some("0") => false,
3160                Some(_) => true,
3161                None => crate::gpu_wgpu::discrete_active(),
3162            } =>
3163        {
3164            crate::gpu_wgpu::dit_block_seg(model, a, segs, x)
3165        }
3166        #[allow(unreachable_patterns)]
3167        _ => false,
3168    }
3169}
3170
3171/// One VAE resnet block for `vae_resnet`: norm/conv weights and the
3172/// channel/shape geometry. `shortcut` is the 1×1 projection (w, b, k)
3173/// when in/out channels differ.
3174pub struct VaeResnetArgs<'a> {
3175    pub groups: usize,
3176    pub ic: usize,
3177    pub oc: usize,
3178    pub h: usize,
3179    pub w: usize,
3180    pub n1w: &'a [f32],
3181    pub n1b: &'a [f32],
3182    pub c1w: &'a [f32],
3183    pub c1b: &'a [f32],
3184    pub c1k: usize,
3185    pub n2w: &'a [f32],
3186    pub n2b: &'a [f32],
3187    pub c2w: &'a [f32],
3188    pub c2b: &'a [f32],
3189    pub c2k: usize,
3190    pub shortcut: Option<(&'a [f32], &'a [f32], usize)>,
3191}
3192
3193/// One whole VAE resnet block on the device (norm+silu → conv ×2 →
3194/// shortcut → add, one command buffer).
3195#[allow(unused_variables)]
3196pub fn vae_resnet(a: &VaeResnetArgs, x: &[f32], out: &mut [f32]) -> bool {
3197    match backend() {
3198        #[cfg(target_os = "macos")]
3199        Backend::Metal => crate::gpu_metal::vae_resnet(a, x, out),
3200        _ => false,
3201    }
3202}
3203
3204/// Nearest-2× upsample fused with the following conv — the small
3205/// pre-upsample image is what crosses the CPU boundary.
3206#[allow(unused_variables, clippy::too_many_arguments)]
3207pub fn vae_upsample_conv(
3208    w: &[f32],
3209    bias: &[f32],
3210    x: &[f32],
3211    ic: usize,
3212    oc: usize,
3213    h: usize,
3214    w_img: usize,
3215    k: usize,
3216    out: &mut [f32],
3217) -> bool {
3218    match backend() {
3219        #[cfg(target_os = "macos")]
3220        Backend::Metal => crate::gpu_metal::vae_upsample_conv(w, bias, x, ic, oc, h, w_img, k, out),
3221        #[cfg(feature = "gpu")]
3222        Backend::Wgpu => crate::gpu_wgpu::vae_upsample_conv(w, bias, x, ic, oc, h, w_img, k, out),
3223        #[allow(unreachable_patterns)]
3224        _ => false,
3225    }
3226}
3227
3228/// VAE conv2d on the device (implicit GEMM — the CPU path pays for a
3229/// multi-GB im2col matrix at high resolutions).
3230#[allow(unused_variables, clippy::too_many_arguments)]
3231pub fn vae_conv2d(
3232    w: &[f32],
3233    bias: &[f32],
3234    x: &[f32],
3235    ic: usize,
3236    oc: usize,
3237    h: usize,
3238    w_img: usize,
3239    k: usize,
3240    out: &mut [f32],
3241) -> bool {
3242    match backend() {
3243        #[cfg(target_os = "macos")]
3244        Backend::Metal => crate::gpu_metal::vae_conv2d(w, bias, x, ic, oc, h, w_img, k, out),
3245        #[cfg(feature = "gpu")]
3246        Backend::Wgpu => crate::gpu_wgpu::vae_conv2d(w, bias, x, ic, oc, h, w_img, k, out),
3247        #[allow(unreachable_patterns)]
3248        _ => false,
3249    }
3250}
3251
3252/// DiT full bidirectional attention on the device (all heads:
3253/// scores GEMM → row softmax → P·V → panel unstack, one command
3254/// buffer). Head-major inputs; out is [n, nh·hd].
3255#[allow(unused_variables, clippy::too_many_arguments)]
3256/// Attention from an interleaved qkv panel, splitting into head-major
3257/// planes ON the device. wgpu only; `false` elsewhere so the caller
3258/// keeps its host repack.
3259#[allow(unused_variables)]
3260#[allow(clippy::too_many_arguments)]
3261/// qkv projection + attention with the panel never leaving the card.
3262/// wgpu only; `false` elsewhere and the caller keeps its host chain.
3263#[allow(clippy::too_many_arguments, unused_variables)]
3264pub fn dit_qkv_attention(
3265    model: &Arc<CmfModel>,
3266    qkv_idx: usize,
3267    xn: &[f32],
3268    n: usize,
3269    hidden: usize,
3270    nh: usize,
3271    hd: usize,
3272    scale: f32,
3273    nr: (&[f32], &[f32], &[f32], f32),
3274    out: &mut [f32],
3275) -> bool {
3276    match backend() {
3277        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3278        Backend::Wgpu => crate::gpu_wgpu::dit_qkv_attention(
3279            model, qkv_idx, xn, n, hidden, nh, hd, scale, nr, out,
3280        ),
3281        #[allow(unreachable_patterns)]
3282        _ => false,
3283    }
3284}
3285
3286/// The whole attention half of a DiT block on the card: qkv GEMM,
3287/// attention, output projection. Only `proj` comes home.
3288#[allow(clippy::too_many_arguments)]
3289pub fn dit_qkv_attn_out(
3290    model: &Arc<CmfModel>,
3291    qkv_idx: usize,
3292    out_idx: usize,
3293    xn: &[f32],
3294    n: usize,
3295    hidden: usize,
3296    nh: usize,
3297    hd: usize,
3298    scale: f32,
3299    nr: (&[f32], &[f32], &[f32], f32),
3300    proj: &mut [f32],
3301) -> bool {
3302    match backend() {
3303        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3304        Backend::Wgpu => crate::gpu_wgpu::dit_qkv_attn_out(
3305            model, qkv_idx, out_idx, xn, n, hidden, nh, hd, scale, nr, proj,
3306        ),
3307        #[allow(unreachable_patterns)]
3308        _ => false,
3309    }
3310}
3311
3312/// The VAE decoder's attention half on the card. Only `proj` returns.
3313#[allow(clippy::too_many_arguments)]
3314pub fn vae_qkv_attn_out(
3315    model: &Arc<CmfModel>,
3316    qkv_idx: usize,
3317    out_idx: usize,
3318    xn: &[f32],
3319    n: usize,
3320    dim: usize,
3321    nh: usize,
3322    hd: usize,
3323    scale: f32,
3324    angles: &[f32],
3325    eps: f32,
3326    qkv_bias: &[f32],
3327    proj: &mut [f32],
3328) -> bool {
3329    match backend() {
3330        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3331        Backend::Wgpu => crate::gpu_wgpu::vae_qkv_attn_out(
3332            model, qkv_idx, out_idx, xn, n, dim, nh, hd, scale, angles, eps, qkv_bias, proj,
3333        ),
3334        #[allow(unreachable_patterns)]
3335        _ => false,
3336    }
3337}
3338
3339#[allow(clippy::too_many_arguments)]
3340pub fn vae_attention_packed(
3341    qkv: &[f32],
3342    nh: usize,
3343    n: usize,
3344    hd: usize,
3345    scale: f32,
3346    angles: &[f32],
3347    eps: f32,
3348    out: &mut [f32],
3349) -> bool {
3350    vae_attention_packed_layout(qkv, nh, n, hd, scale, angles, eps, out, 1)
3351}
3352
3353#[allow(clippy::too_many_arguments)]
3354pub fn vae_attention_packed_layout(
3355    qkv: &[f32],
3356    nh: usize,
3357    n: usize,
3358    hd: usize,
3359    scale: f32,
3360    angles: &[f32],
3361    eps: f32,
3362    out: &mut [f32],
3363    layout: u32,
3364) -> bool {
3365    match backend() {
3366        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3367        Backend::Wgpu => crate::gpu_wgpu::vae_attention_packed_layout(
3368            qkv, nh, n, hd, scale, angles, eps, out, layout,
3369        ),
3370        #[allow(unreachable_patterns)]
3371        _ => false,
3372    }
3373}
3374
3375#[allow(clippy::too_many_arguments)]
3376pub fn dit_split_only(
3377    qkv: &[f32],
3378    nh: usize,
3379    n: usize,
3380    hd: usize,
3381    layout: u32,
3382    norm: Option<(&[f32], f32)>,
3383    out_q: &mut [f32],
3384) -> bool {
3385    match backend() {
3386        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3387        Backend::Wgpu => crate::gpu_wgpu::dit_split_only(qkv, nh, n, hd, layout, norm, out_q),
3388        #[allow(unreachable_patterns)]
3389        _ => false,
3390    }
3391}
3392
3393/// The backend's f32 NT GEMM: `y[n×m] = x[n×k] · wᵀ[m×k]`. Tensor
3394/// cores where the card has them. Refuses under `CMF_BAKE_GPU=0` or
3395/// strict f32, and for jobs below n·k·m = 4M, where the round trip
3396/// costs more than the arithmetic saves.
3397/// `gemm_nt_f32` whose `w` is known to change every call (an
3398/// accumulation over fresh activations, not a weight): it skips the
3399/// resident ledger and its per-call fingerprint of the whole operand.
3400pub fn gemm_nt_f32_transient(
3401    x: &[f32],
3402    w: &[f32],
3403    y: &mut [f32],
3404    n: usize,
3405    k: usize,
3406    m: usize,
3407) -> bool {
3408    match backend() {
3409        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3410        Backend::Wgpu => crate::gpu_wgpu::gemm_nt_f32_transient(x, w, y, n, k, m),
3411        #[allow(unreachable_patterns)]
3412        _ => false,
3413    }
3414}
3415
3416pub fn gemm_nt_f32(x: &[f32], w: &[f32], y: &mut [f32], n: usize, k: usize, m: usize) -> bool {
3417    match backend() {
3418        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3419        Backend::Wgpu => crate::gpu_wgpu::gemm_nt_f32(x, w, y, n, k, m),
3420        #[allow(unreachable_patterns)]
3421        _ => false,
3422    }
3423}
3424
3425/// Music-3's FFN chain resident on the device — two GEMMs and the GLU
3426/// between them with no host round trip. `false` = refused, host runs.
3427#[allow(clippy::too_many_arguments)]
3428pub fn music3_ffn(
3429    model: &std::sync::Arc<CmfModel>,
3430    idx_in: usize,
3431    idx_out: usize,
3432    h: &[f32],
3433    bias_in: &[f32],
3434    n: usize,
3435    hs: usize,
3436    inter: usize,
3437    out: &mut [f32],
3438) -> bool {
3439    match backend() {
3440        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3441        Backend::Wgpu => {
3442            crate::gpu_wgpu::music3_ffn(model, idx_in, idx_out, h, bias_in, n, hs, inter, out)
3443        }
3444        #[allow(unreachable_patterns)]
3445        _ => false,
3446    }
3447}
3448
3449/// A 1D convolution as a GEMM whose column matrix is expanded on the
3450/// device instead of being built, transposed and uploaded by the host.
3451/// `yt` comes back `[out_n x oc]`. `false` = refused, caller runs host.
3452#[allow(clippy::too_many_arguments)]
3453pub fn conv1d_gemm(
3454    x: &[f32],
3455    w: &[f32],
3456    ic: usize,
3457    oc: usize,
3458    n: usize,
3459    k: usize,
3460    pad: usize,
3461    dil: usize,
3462    out_n: usize,
3463    yt: &mut [f32],
3464) -> bool {
3465    match backend() {
3466        #[cfg(target_os = "macos")]
3467        Backend::Metal => crate::gpu_metal::conv1d_gemm(x, w, ic, oc, n, k, pad, dil, out_n, yt),
3468        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3469        Backend::Wgpu => crate::gpu_wgpu::conv1d_gemm(x, w, ic, oc, n, k, pad, dil, out_n, yt),
3470        #[allow(unreachable_patterns)]
3471        _ => false,
3472    }
3473}
3474
3475/// The convolution as a GEMM on the matrix units. `false` = refused.
3476#[allow(clippy::too_many_arguments)]
3477pub fn vae_conv2d_coop(
3478    w: &[f32],
3479    bias: Option<&[f32]>,
3480    x: &[f32],
3481    ic: usize,
3482    oc: usize,
3483    h: usize,
3484    wi: usize,
3485    k: usize,
3486    out: &mut [f32],
3487) -> bool {
3488    match backend() {
3489        #[cfg(all(feature = "gpu", not(target_os = "macos")))]
3490        Backend::Wgpu => crate::gpu_wgpu::vae_conv2d_coop(w, bias, x, ic, oc, h, wi, k, out),
3491        #[allow(unreachable_patterns)]
3492        _ => false,
3493    }
3494}
3495
3496pub fn dit_attention_packed(
3497    qkv: &[f32],
3498    nh: usize,
3499    n: usize,
3500    hd: usize,
3501    scale: f32,
3502    // (rope angles, q norm weights, k norm weights, eps) when the device
3503    // should apply qk-norm and RoPE itself; None when the host already did.
3504    nr: Option<(&[f32], &[f32], &[f32], f32)>,
3505    out: &mut [f32],
3506) -> bool {
3507    match backend() {
3508        // wgpu carries the only implementation, and it is not
3509        // platform-specific: `CMF_GPU=wgpu` on macOS runs it over Metal
3510        // like anywhere else. It used to be compiled out here on macOS,
3511        // which made the call a silent `false` — and the caller's
3512        // `assert!` turned that refusal into a panic on every
3513        // `cortiq animate` this platform ever ran.
3514        #[cfg(feature = "gpu")]
3515        Backend::Wgpu => crate::gpu_wgpu::dit_attention_packed(qkv, nh, n, hd, scale, nr, out),
3516        #[allow(unreachable_patterns)]
3517        _ => false,
3518    }
3519}
3520
3521/// Whether `dit_attention_packed` has an implementation on the backend
3522/// that is actually selected.
3523///
3524/// The caller has to know BEFORE it skips the host qk-norm: deferring
3525/// the norm to a device that then refuses leaves q/k unnormalized with
3526/// no way back. Native Metal has no packed kernel, so on macOS this is
3527/// false unless `CMF_GPU=wgpu` picked the other backend.
3528pub fn dit_attention_packed_available() -> bool {
3529    #[allow(unreachable_patterns)]
3530    match backend() {
3531        #[cfg(feature = "gpu")]
3532        Backend::Wgpu => crate::gpu_wgpu::dit_attention_packed_ready(),
3533        _ => false,
3534    }
3535}
3536
3537pub fn dit_attention(
3538    qh: &[f32],
3539    kh: &[f32],
3540    vh: &[f32],
3541    nh: usize,
3542    nkv: usize,
3543    n: usize,
3544    hd: usize,
3545    scale: f32,
3546    out: &mut [f32],
3547) -> bool {
3548    match backend() {
3549        #[cfg(target_os = "macos")]
3550        Backend::Metal => crate::gpu_metal::dit_attention(qh, kh, vh, nh, nkv, n, hd, scale, out),
3551        #[cfg(feature = "gpu")]
3552        Backend::Wgpu => crate::gpu_wgpu::dit_attention(qh, kh, vh, nh, nkv, n, hd, scale, out),
3553        #[allow(unreachable_patterns)]
3554        _ => false,
3555    }
3556}
3557
3558/// Batched q4t GEMM on the device (imagegen DiT prefill shapes).
3559/// Metal: q4t_mul_mm decodes the mmap-resident tiles inside the
3560/// GEMM's K loop. wgpu (Vulkan/DX12 → NVIDIA/AMD/Intel/Adreno/Mali):
3561/// the register-blocked WGSL twin, weights cached in VRAM.
3562#[allow(unused_variables)]
3563pub fn q4tp_matmat(
3564    model: &Arc<CmfModel>,
3565    idx: usize,
3566    xs: &[f32],
3567    b: usize,
3568    rows: usize,
3569    cols: usize,
3570    out: &mut [f32],
3571) -> bool {
3572    match backend() {
3573        #[cfg(target_os = "macos")]
3574        Backend::Metal => crate::gpu_metal::q4tp_matmat(model, idx, xs, b, rows, cols, out),
3575        #[cfg(feature = "gpu")]
3576        Backend::Wgpu => crate::gpu_wgpu::q4tp_matmat(model, idx, xs, b, rows, cols, out),
3577        #[allow(unreachable_patterns)]
3578        _ => false,
3579    }
3580}
3581
3582/// The same over a two-bit weight plane. Native Metal uses the dedicated
3583/// q2tp tile; unsupported shapes return false and preserve the host fallback.
3584pub fn q2tp_matmat(
3585    model: &Arc<CmfModel>,
3586    idx: usize,
3587    xs: &[f32],
3588    b: usize,
3589    rows: usize,
3590    cols: usize,
3591    out: &mut [f32],
3592) -> bool {
3593    match backend() {
3594        #[cfg(target_os = "macos")]
3595        Backend::Metal => crate::gpu_metal::q2tp_matmat(model, idx, xs, b, rows, cols, out),
3596        #[cfg(feature = "gpu")]
3597        Backend::Wgpu => crate::gpu_wgpu::q2tp_matmat(model, idx, xs, b, rows, cols, out),
3598        #[allow(unreachable_patterns)]
3599        _ => false,
3600    }
3601}
3602
3603/// Descriptor-aware q2tp GEMM. The affine center is selected only for a
3604/// validated q2tp_affine target; the raw dtype16 payload remains unchanged.
3605pub fn q2tp_affine_matmat(
3606    model: &Arc<CmfModel>,
3607    idx: usize,
3608    xs: &[f32],
3609    b: usize,
3610    rows: usize,
3611    cols: usize,
3612    out: &mut [f32],
3613) -> bool {
3614    match backend() {
3615        #[cfg(target_os = "macos")]
3616        Backend::Metal => crate::gpu_metal::q2tp_affine_matmat(model, idx, xs, b, rows, cols, out),
3617        #[cfg(feature = "gpu")]
3618        Backend::Wgpu => crate::gpu_wgpu::q2tp_affine_matmat(model, idx, xs, b, rows, cols, out),
3619        #[allow(unreachable_patterns)]
3620        _ => false,
3621    }
3622}
3623
3624/// Single-token q2tp matvec through the ordinary (center=1.5) WGSL kernel.
3625pub fn q2tp_matvec(
3626    model: &Arc<CmfModel>,
3627    idx: usize,
3628    xs: &[f32],
3629    rows: usize,
3630    cols: usize,
3631    out: &mut [f32],
3632) -> bool {
3633    match backend() {
3634        #[cfg(target_os = "macos")]
3635        Backend::Metal => crate::gpu_metal::q2tp_matvec(model, idx, xs, rows, cols, out),
3636        #[cfg(feature = "gpu")]
3637        Backend::Wgpu => crate::gpu_wgpu::q2tp_matvec(model, idx, xs, rows, cols, out),
3638        #[allow(unreachable_patterns)]
3639        _ => false,
3640    }
3641}
3642
3643/// Single-token q2tp matvec with the explicit affine center=1 descriptor
3644/// operator. This is kept separate from ordinary q2tp to make accidental
3645/// center changes impossible at a call site.
3646pub fn q2tp_affine_matvec(
3647    model: &Arc<CmfModel>,
3648    idx: usize,
3649    xs: &[f32],
3650    rows: usize,
3651    cols: usize,
3652    out: &mut [f32],
3653) -> bool {
3654    match backend() {
3655        #[cfg(target_os = "macos")]
3656        Backend::Metal => crate::gpu_metal::q2tp_affine_matvec(model, idx, xs, rows, cols, out),
3657        #[cfg(feature = "gpu")]
3658        Backend::Wgpu => crate::gpu_wgpu::q2tp_affine_matvec(model, idx, xs, rows, cols, out),
3659        #[allow(unreachable_patterns)]
3660        _ => false,
3661    }
3662}
3663
3664/// Single-token q4tp matvec on the device — the lm_head class. Through the
3665/// DEDICATED matvec kernel: the batched GEMM at b=1 measured 11.73 ms
3666/// against the host's 9.51 on the release head, so the route that was
3667/// supposed to save eleven milliseconds a token lost its own probe instead.
3668pub fn q4tp_matvec(
3669    model: &Arc<CmfModel>,
3670    idx: usize,
3671    xs: &[f32],
3672    rows: usize,
3673    cols: usize,
3674    out: &mut [f32],
3675) -> bool {
3676    match backend() {
3677        #[cfg(target_os = "macos")]
3678        Backend::Metal => crate::gpu_metal::q4tp_matvec_for_test(model, idx, xs, rows, cols, out),
3679        #[cfg(feature = "gpu")]
3680        Backend::Wgpu => crate::gpu_wgpu::q4tp_matvec(model, idx, xs, rows, cols, out),
3681        #[allow(unreachable_patterns)]
3682        _ => false,
3683    }
3684}
3685
3686/// Single-token q4_tiled matvec on the device — the lm_head class (a
3687/// q4t checkpoint's head is its biggest host matvec, exactly like the
3688/// q4tp twin above). wgpu holds q4t_mv pipelines only inside the graph
3689/// encoder — the standalone arm stays an honest refusal until a
3690/// discrete-GPU q4t model reaches the bench.
3691pub fn q4t_matvec(
3692    model: &Arc<CmfModel>,
3693    idx: usize,
3694    xs: &[f32],
3695    rows: usize,
3696    cols: usize,
3697    out: &mut [f32],
3698) -> bool {
3699    match backend() {
3700        #[cfg(target_os = "macos")]
3701        Backend::Metal => crate::gpu_metal::q4t_matvec_for_test(model, idx, xs, rows, cols, out),
3702        #[allow(unreachable_patterns)]
3703        _ => false,
3704    }
3705}
3706
3707pub fn q4t_matmat(
3708    model: &Arc<CmfModel>,
3709    idx: usize,
3710    xs: &[f32],
3711    b: usize,
3712    rows: usize,
3713    cols: usize,
3714    out: &mut [f32],
3715) -> bool {
3716    match backend() {
3717        #[cfg(target_os = "macos")]
3718        Backend::Metal => crate::gpu_metal::q4t_matmat(model, idx, xs, b, rows, cols, out),
3719        #[cfg(feature = "gpu")]
3720        Backend::Wgpu => crate::gpu_wgpu::q4t_matmat(model, idx, xs, b, rows, cols, out),
3721        #[allow(unreachable_patterns)]
3722        _ => false,
3723    }
3724}
3725
3726/// Whole-block token-graph types re-exported from the Metal backend.
3727#[cfg(target_os = "macos")]
3728pub use crate::gpu_metal::{
3729    AttnDeviceParams, AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GpuMoe, GraphDims, MetalFfn,
3730    O1AttnParams, TokenGraph, kv_mirror_drop, kv_mirror_read_last, kv_mirror_take_imp,
3731};
3732
3733/// A BLOCK of consecutive q1 GDN layers in one submission (Metal only).
3734#[cfg(target_os = "macos")]
3735pub fn gdn_block(
3736    model: &Arc<CmfModel>,
3737    layers: &[GdnGpuLayer],
3738    states: &mut [&mut [f32]],
3739    cfg: &GdnGpuCfg,
3740    h: &mut [f32],
3741) -> bool {
3742    match backend() {
3743        Backend::Metal => crate::gpu_metal::gdn_block(model, layers, states, cfg, h),
3744        _ => false,
3745    }
3746}
3747
3748/// A layer's MoE-FFN in one submission (amortizing the dispatch cost).
3749#[allow(unused_variables)]
3750pub fn moe_block(model: &Arc<CmfModel>, jobs: &[MoeJob], out: &mut [f32]) -> bool {
3751    match backend() {
3752        #[cfg(target_os = "macos")]
3753        Backend::Metal => crate::gpu_metal::moe_block(model, jobs, out),
3754        #[cfg(feature = "gpu")]
3755        Backend::Wgpu => crate::gpu_wgpu::moe_block(model, jobs, out),
3756        Backend::None => false,
3757    }
3758}
3759
3760/// Independent matvecs of one input in a single submission (GDN projections).
3761#[allow(unused_variables)]
3762pub fn matvec_batch(model: &Arc<CmfModel>, jobs: &[BatchJob], out: &mut [&mut [f32]]) -> bool {
3763    match backend() {
3764        #[cfg(target_os = "macos")]
3765        Backend::Metal => crate::gpu_metal::matvec_batch(model, jobs, out),
3766        #[cfg(feature = "gpu")]
3767        Backend::Wgpu => crate::gpu_wgpu::matvec_batch(model, jobs, out),
3768        Backend::None => false,
3769    }
3770}
3771
3772// ── Whole-token wgpu graph race (generation granularity) ─────────────
3773// On integrated/mobile adapters the graph is neither trusted nor banned
3774// a priori — it RACES the normal path: generations alternate arms (the
3775// normal path first — known-good UX — then the graph), per-token wall
3776// times accumulate per arm, and once both arms have enough steady
3777// samples the faster one wins for the process. Arm switches happen ONLY
3778// at generation boundaries (`kv_cache.clear()` resets state), so the
3779// device KV mirror and the CPU cache never diverge mid-sequence. The
3780// single exception is the first-token bail: the very first decode token
3781// of a graph generation may be discarded and recomputed on the CPU
3782// path (the prompt KV is CPU-owned at that point, so this is safe) —
3783// a tiled mobile GPU that drains its pipeline at every barrier turns
3784// the ~300-dispatch graph into seconds per token (field report: 0.2
3785// tok/s vs 15 on the CPU), and one token is all it takes to see that.
3786static GRAPH_RACE_STATE: AtomicU8 = AtomicU8::new(0); // 0 racing, 1 graph won, 2 normal won
3787static GRAPH_RACE_FLIP: AtomicU32 = AtomicU32::new(0);
3788static GRAPH_RACE_ARM_GRAPH: AtomicU8 = AtomicU8::new(0); // this generation's arm
3789static GRAPH_RACE_TOK: AtomicU32 = AtomicU32::new(0); // token index within the generation
3790static GRAPH_NS: [AtomicU64; 2] = [AtomicU64::new(0), AtomicU64::new(0)]; // [normal, graph]
3791static GRAPH_N: [AtomicU32; 2] = [AtomicU32::new(0), AtomicU32::new(0)];
3792
3793/// Steady per-token samples per arm before the race decides.
3794const GRAPH_RACE_SAMPLES: u32 = 4;
3795
3796// A graph that cannot be built for THIS model will never build: the
3797// refusal is a property of the weights, not of the moment. Retrying it
3798// per token is not free — the builder walks every layer and asks each
3799// tensor for a graph view before giving up at layer 0 — and on an
3800// Adreno 642L that retry cost 3x: forcing the graph on a model it
3801// refuses measured 0.3 tok/s against 0.905 for the per-op path it falls
3802// back to. Remembered once, the fallback runs at its own speed.
3803//
3804// The verdict is kept PER PIPELINE (`Pipeline::graph_refused` /
3805// `mark_graph_refused`), not in a process-wide flag: several pipelines
3806// of one file share a process in `serve` (backbone slots + skill lanes
3807// loaded mid-traffic), and a refusal in one lane — or the reset a new
3808// lane used to issue — must not move another lane's running sequence
3809// between the device and the host (R4/NF-2). Callers must NOT report
3810// transient refusals (an unsealed o1 state during prefill, a softcap).
3811
3812/// Called at every generation start (fresh KV). Applies a pending
3813/// verdict and picks this generation's arm while racing.
3814pub fn graph_race_begin_generation() {
3815    // One generation has now compiled whatever this model needs; keep it
3816    // for the next process. Once per run: the blob does not grow after
3817    // the pipelines exist, and the write is megabytes against the ~200 s
3818    // of compiling it saves on the device that needed this.
3819    #[cfg(feature = "gpu")]
3820    {
3821        // Save once, at the start of the SECOND generation: the first
3822        // has dispatched, so there is something to keep, and nothing is
3823        // saved before any work (the driver compiles at first use, not
3824        // at pipeline creation — the context comes up in 1.5 s while the
3825        // compiling costs minutes).
3826        //
3827        // Flushing again on 4, 8, 16 … was tried on the theory that a
3828        // chat turn compiles shapes the first one did not. It buys
3829        // nothing: a fresh app process still spent 49.0 s, then 58.7,
3830        // then 61.3 on its first answer with the backoff in place. One
3831        // flush it is.
3832        static FLUSHED: std::sync::Once = std::sync::Once::new();
3833        static FIRST: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(true);
3834        if FIRST.swap(false, Ordering::Relaxed) {
3835            // Nothing dispatched yet.
3836        } else {
3837            FLUSHED.call_once(crate::gpu_wgpu::pipeline_cache_flush);
3838        }
3839    }
3840    GRAPH_RACE_TOK.store(0, Ordering::Relaxed);
3841    if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
3842        return;
3843    }
3844    let (gn, cn) = (
3845        GRAPH_N[1].load(Ordering::Relaxed),
3846        GRAPH_N[0].load(Ordering::Relaxed),
3847    );
3848    if gn >= GRAPH_RACE_SAMPLES && cn >= GRAPH_RACE_SAMPLES {
3849        let g_avg = GRAPH_NS[1].load(Ordering::Relaxed) / gn as u64;
3850        let c_avg = GRAPH_NS[0].load(Ordering::Relaxed) / cn as u64;
3851        let verdict = if g_avg < c_avg { 1 } else { 2 };
3852        GRAPH_RACE_STATE.store(verdict, Ordering::Relaxed);
3853        tracing::info!(
3854            "wgpu graph race: graph {:.2} ms/tok vs normal {:.2} ms/tok -> {}",
3855            g_avg as f64 / 1e6,
3856            c_avg as f64 / 1e6,
3857            if verdict == 1 { "graph" } else { "normal path" }
3858        );
3859        return;
3860    }
3861    let flip = GRAPH_RACE_FLIP.fetch_add(1, Ordering::Relaxed);
3862    GRAPH_RACE_ARM_GRAPH.store((flip % 2 == 1) as u8, Ordering::Relaxed);
3863}
3864
3865/// Should this decode token try the graph? `trusted` (discrete adapter,
3866/// explicit env, or a GDN hybrid whose state lives on the device) skips
3867/// the race entirely.
3868pub fn graph_race_use_graph(trusted: bool) -> bool {
3869    if trusted {
3870        return true;
3871    }
3872    match GRAPH_RACE_STATE.load(Ordering::Relaxed) {
3873        1 => true,
3874        2 => false,
3875        _ => GRAPH_RACE_ARM_GRAPH.load(Ordering::Relaxed) == 1,
3876    }
3877}
3878
3879/// First decode token of a racing graph generation: hopeless already?
3880/// (>4x the normal path's per-token average AND over a second.) Settles
3881/// the race immediately; the caller discards the graph result and
3882/// recomputes this token on the normal path.
3883pub fn graph_race_first_token_hopeless(dur: std::time::Duration) -> bool {
3884    if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
3885        return false;
3886    }
3887    let first = GRAPH_RACE_TOK.load(Ordering::Relaxed) == 0;
3888    let cn = GRAPH_N[0].load(Ordering::Relaxed);
3889    if !first || cn == 0 {
3890        return false;
3891    }
3892    let c_avg = GRAPH_NS[0].load(Ordering::Relaxed) / cn as u64;
3893    let ns = dur.as_nanos() as u64;
3894    if ns > 1_000_000_000 && ns > 4 * c_avg {
3895        GRAPH_RACE_STATE.store(2, Ordering::Relaxed);
3896        tracing::info!(
3897            "wgpu graph race: first graph token {:.0} ms vs normal {:.2} ms/tok — hopeless, normal path wins",
3898            ns as f64 / 1e6,
3899            c_avg as f64 / 1e6
3900        );
3901        return true;
3902    }
3903    false
3904}
3905
3906/// Record one decode-token wall time for the racing arm. The first
3907/// token of each generation is discarded (KV-mirror upload / cold
3908/// caches on the graph arm; cold mmap on the normal arm).
3909pub fn graph_race_record(used_graph: bool, dur: std::time::Duration) {
3910    if GRAPH_RACE_STATE.load(Ordering::Relaxed) != 0 {
3911        return;
3912    }
3913    let tok = GRAPH_RACE_TOK.fetch_add(1, Ordering::Relaxed);
3914    if tok == 0 {
3915        return;
3916    }
3917    let i = used_graph as usize;
3918    GRAPH_NS[i].fetch_add(dur.as_nanos() as u64, Ordering::Relaxed);
3919    GRAPH_N[i].fetch_add(1, Ordering::Relaxed);
3920}
3921
3922/// Bounded-cost content fingerprint for the backends' pointer-keyed device
3923/// caches: FNV over the whole slice up to 4 KiB, over 64 spread 64-byte
3924/// windows (plus the length) above. An address-keyed hit must also prove
3925/// the bytes are still the ones it uploaded — the allocator reuses heap
3926/// and mmap addresses freely, so a reloaded model or a re-dequantized
3927/// layer lands where the old bytes were — and sampling keeps that proof at
3928/// ~a microsecond even for a 126 MB matrix. Real replacements (another
3929/// model's tensor, an Adam-updated master) differ densely, so a 4 KiB
3930/// spread cannot miss them.
3931pub(crate) fn fp_bytes(data: &[u8]) -> u64 {
3932    #[inline]
3933    fn fnv(mut h: u64, bytes: &[u8]) -> u64 {
3934        let (chunks, tail) = bytes.split_at(bytes.len() & !7);
3935        for c in chunks.chunks_exact(8) {
3936            h ^= u64::from_le_bytes(c.try_into().unwrap());
3937            h = h.wrapping_mul(0x100_0000_01b3);
3938        }
3939        for &b in tail {
3940            h ^= b as u64;
3941            h = h.wrapping_mul(0x100_0000_01b3);
3942        }
3943        h
3944    }
3945    let mut h = 0xcbf2_9ce4_8422_2325u64 ^ (data.len() as u64);
3946    if data.len() <= 4096 {
3947        return fnv(h, data);
3948    }
3949    let step = (data.len() - 64) / 63;
3950    for i in 0..64 {
3951        h = fnv(h, &data[i * step..i * step + 64]);
3952    }
3953    h
3954}
3955
3956/// `fp_bytes` over an f32 slice without a bytemuck dependency (the Metal
3957/// backend builds with no GPU feature flags).
3958pub(crate) fn fp_f32(data: &[f32]) -> u64 {
3959    let bytes = unsafe { std::slice::from_raw_parts(data.as_ptr() as *const u8, data.len() * 4) };
3960    fp_bytes(bytes)
3961}
3962
3963#[cfg(test)]
3964mod fp_tests {
3965    use super::fp_bytes;
3966
3967    /// The pointer-keyed caches survive on `fp_bytes` telling two different
3968    /// tensors apart at a reused address. Its sampling must therefore see a
3969    /// change ANYWHERE — head, tail, and the stretches between windows are
3970    /// the places a cheaper hash would go blind.
3971    #[test]
3972    fn fp_bytes_sees_a_change_anywhere_in_a_sampled_slice() {
3973        let n = 1 << 20; // 1 MiB — far above the 4 KiB full-hash threshold
3974        let base: Vec<u8> = (0..n).map(|i| (i * 31 + 7) as u8).collect();
3975        let h0 = fp_bytes(&base);
3976        assert_eq!(h0, fp_bytes(&base), "fingerprint must be deterministic");
3977        // A DENSE change (every requantized/redequantized tensor is one)
3978        // must flip the fingerprint no matter how the windows fall.
3979        let mut dense = base.clone();
3980        for b in dense.iter_mut() {
3981            *b = b.wrapping_add(1);
3982        }
3983        assert_ne!(
3984            h0,
3985            fp_bytes(&dense),
3986            "a fully different tensor slipped through"
3987        );
3988        // Length participates: the same prefix at a shorter length is a
3989        // different key AND a different fingerprint.
3990        assert_ne!(h0, fp_bytes(&base[..n - 64]));
3991        // Below the threshold the hash is exact: a single flipped byte in
3992        // a norm-sized vector must be seen.
3993        let mut small = vec![3u8; 4096];
3994        let hs = fp_bytes(&small);
3995        small[2048] ^= 1;
3996        assert_ne!(hs, fp_bytes(&small), "full hash missed a one-byte change");
3997        // And the sampled windows land within bounds on awkward sizes.
3998        for n in [4097usize, 5000, 64 * 64, 1 << 16] {
3999            let v = vec![9u8; n];
4000            let _ = fp_bytes(&v); // must not panic on window math
4001        }
4002    }
4003}
4004
4005/// Hand the card back after a bake: drop its resident weights, planes and
4006/// pools so the ordinary engine (the runtime gate, a serve that follows)
4007/// starts from a clean budget. No-op off the wgpu backend.
4008pub fn bake_release() {
4009    #[cfg(feature = "gpu")]
4010    crate::gpu_wgpu::bake_release();
4011}
4012
4013/// Strict-f32 for the bake's GEMMs (phase A mask training): the mask
4014/// selects neurons by a gradient signal, and f16 operand rounding on
4015/// that signal closes the wrong ones. No-op off the wgpu backend.
4016pub fn bake_precision_strict(on: bool) {
4017    #[cfg(feature = "gpu")]
4018    crate::gpu_wgpu::bake_precision_strict(on);
4019    #[cfg(not(feature = "gpu"))]
4020    let _ = on;
4021}
4022
4023/// CMF_GRAPH_HOSTPROF=1: how a graph token's wall splits between the
4024/// host encoding the command stream and the tail the GPU still owes
4025/// after encode. Fifteen GPU-side suspects measured null while the
4026/// bench counted 17.7k allocations a token — this is the instrument
4027/// that says whether the thief was on the host all along.
4028pub fn hostprof_encode_done(t0: std::time::Instant) {
4029    use std::sync::atomic::{AtomicU64, Ordering};
4030    static ENC: AtomicU64 = AtomicU64::new(0);
4031    static N: AtomicU64 = AtomicU64::new(0);
4032    if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
4033        return;
4034    }
4035    ENC.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
4036    let n = N.fetch_add(1, Ordering::Relaxed) + 1;
4037    if n % 100 == 0 {
4038        eprintln!(
4039            "hostprof: encode {:.2} ms/token over {n} tokens",
4040            ENC.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
4041        );
4042    }
4043}
4044
4045pub fn hostprof_total(t0: std::time::Instant) {
4046    use std::sync::atomic::{AtomicU64, Ordering};
4047    static TOT: AtomicU64 = AtomicU64::new(0);
4048    static N: AtomicU64 = AtomicU64::new(0);
4049    if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
4050        return;
4051    }
4052    TOT.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
4053    let n = N.fetch_add(1, Ordering::Relaxed) + 1;
4054    if n % 100 == 0 {
4055        eprintln!(
4056            "hostprof: total {:.2} ms/token over {n} tokens",
4057            TOT.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
4058        );
4059    }
4060}
4061
4062/// Per-stage host-encode accumulator for the Metal token loop
4063/// (CMF_GRAPH_HOSTPROF=1). Stage 0 = GDN-run encode; everything else
4064/// falls out by subtraction from hostprof's encode total.
4065pub fn stageprof(stage: u32, dt: std::time::Duration) {
4066    use std::sync::atomic::{AtomicU64, Ordering};
4067    static NS: [AtomicU64; 4] = [
4068        AtomicU64::new(0),
4069        AtomicU64::new(0),
4070        AtomicU64::new(0),
4071        AtomicU64::new(0),
4072    ];
4073    static N: AtomicU64 = AtomicU64::new(0);
4074    if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() != Ok("1") {
4075        return;
4076    }
4077    NS[stage as usize % 4].fetch_add(dt.as_nanos() as u64, Ordering::Relaxed);
4078    if stage == 1 {
4079        let n = N.fetch_add(1, Ordering::Relaxed) + 1;
4080        if n % 200 == 0 {
4081            eprintln!(
4082                "stageprof: planning {:.2} ms/tok | gdn-item {:.2} ms/tok | attn-item {:.2} ms/tok ({n} tok)",
4083                NS[1].load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
4084                NS[2].load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
4085                NS[3].load(Ordering::Relaxed) as f64 / n as f64 / 1e6
4086            );
4087        }
4088    }
4089}
4090
4091/// Active weight bytes dispatched so far (Metal decode path); 0 where
4092/// the backend does not count. The honest floor's numerator.
4093pub fn weight_bytes_dispatched() -> u64 {
4094    let mut total = 0u64;
4095    #[cfg(target_os = "macos")]
4096    {
4097        total += crate::gpu_metal::WEIGHT_BYTES.load(std::sync::atomic::Ordering::Relaxed);
4098    }
4099    #[cfg(feature = "gpu")]
4100    {
4101        total += crate::gpu_wgpu::WEIGHT_BYTES.load(std::sync::atomic::Ordering::Relaxed);
4102    }
4103    total
4104}
4105
4106/// The per-stage split of `weight_bytes_dispatched`:
4107/// [misc, dense-ffn, moe, attn, gdn, head].
4108pub fn weight_bytes_by() -> [u64; 6] {
4109    #[cfg(target_os = "macos")]
4110    {
4111        let mut o = [0u64; 6];
4112        for (i, a) in crate::gpu_metal::WEIGHT_BYTES_BY.iter().enumerate() {
4113            o[i] = a.load(std::sync::atomic::Ordering::Relaxed);
4114        }
4115        return o;
4116    }
4117    #[allow(unreachable_code)]
4118    [0; 6]
4119}
4120
4121#[cfg(test)]
4122mod f16_guard_tests {
4123    use super::*;
4124
4125    /// The f16 guard's maximum: the largest magnitude through the lanes and
4126    /// the remainder, +inf for an inf or a NaN anywhere.
4127    #[test]
4128    fn abs_max_or_inf_is_the_magnitude_maximum() {
4129        let mut xs: Vec<f32> = (0..37).map(|i| (i as f32 - 18.0) * 0.5).collect();
4130        assert_eq!(abs_max_or_inf(&xs), 9.0);
4131        xs[36] = -70000.0; // the remainder
4132        assert_eq!(abs_max_or_inf(&xs), 70000.0);
4133        xs[3] = f32::NAN; // a lane
4134        assert_eq!(abs_max_or_inf(&xs), f32::INFINITY);
4135        xs[3] = 0.0;
4136        xs[35] = f32::NEG_INFINITY;
4137        assert_eq!(abs_max_or_inf(&xs), f32::INFINITY);
4138        assert_eq!(abs_max_or_inf(&[]), 0.0);
4139        assert_eq!(abs_max_or_inf(&[-0.0, -1e-30]), 1e-30);
4140    }
4141}
4142
4143#[cfg(test)]
4144mod probe_warmup_tests {
4145    use super::*;
4146    use std::time::Duration;
4147
4148    fn ms(v: f64) -> Duration {
4149        Duration::from_nanos((v * 1e6) as u64)
4150    }
4151
4152    /// The bug this pins, measured on an A100: the first device call for
4153    /// a class compiles its pipeline, was timed at 117.01 ms against the
4154    /// host's 3.19, and sent `gemm-nt` to the CPU for the whole process —
4155    /// which ran a 27B bake on 2.6 cores with the card idle.
4156    #[test]
4157    fn one_cold_first_sample_does_not_lose_the_class() {
4158        let p = Probe::new();
4159        // First device sample is the pipeline build. Then the truth.
4160        probe_record_into(&p, "gemm-nt", None, true, ms(117.01));
4161        probe_record_into(&p, "gemm-nt", None, true, ms(1.1));
4162        probe_record_into(&p, "gemm-nt", None, true, ms(1.0));
4163        probe_record_into(&p, "gemm-nt", None, false, ms(3.19));
4164        probe_record_into(&p, "gemm-nt", None, false, ms(3.20));
4165        assert_eq!(
4166            p.state.load(Ordering::Relaxed),
4167            1,
4168            "the device is 3x faster once warm and must win"
4169        );
4170    }
4171
4172    /// The warm-up must not become a way to never decide, and must not
4173    /// underflow: a blind decrement at zero wraps a u32 to its maximum
4174    /// and mutes the arm for the life of the process.
4175    #[test]
4176    fn the_warmup_is_spent_once_and_never_underflows() {
4177        let p = Probe::new();
4178        for _ in 0..8 {
4179            probe_record_into(&p, "matmat", None, true, ms(10.0));
4180        }
4181        assert_eq!(p.gpu_burn.load(Ordering::Relaxed), 0, "spent, not wrapped");
4182        assert_eq!(
4183            p.gpu_n.load(Ordering::Relaxed),
4184            7,
4185            "one sample burned, the rest counted"
4186        );
4187    }
4188
4189    /// A device path that always refuses records no timing, so without
4190    /// counting the refusals the class can never reach a verdict. On an
4191    /// M4 with LFM2.5-2.6B `ffn` was still undecided after 9000 calls,
4192    /// alternating arms and paying a failed device attempt on half of
4193    /// them.
4194    #[test]
4195    fn a_class_whose_device_always_declines_settles_on_the_host() {
4196        let _probe_guard = probe_test_guard();
4197        // A class no other test in this file touches: `probe_note_decline`
4198        // works on the process-wide probes by design, and the tests in
4199        // this binary share them.
4200        let c = OpClass::MatmatWide;
4201        let p = &PROBES[c as usize];
4202        p.state.store(0, Ordering::Relaxed);
4203        p.declines.store(0, Ordering::Relaxed);
4204        for _ in 0..(PROBE_DECLINE_LIMIT - 1) {
4205            probe_note_decline(c);
4206        }
4207        assert_eq!(
4208            p.state.load(Ordering::Relaxed),
4209            0,
4210            "one short of the limit is still a question, not an answer"
4211        );
4212        probe_note_decline(c);
4213        assert_eq!(p.state.load(Ordering::Relaxed), 2, "settled on the host");
4214        assert!(matches!(probe_arm(c), ProbeArm::Cpu));
4215        p.state.store(0, Ordering::Relaxed);
4216        p.declines.store(0, Ordering::Relaxed);
4217    }
4218
4219    /// A genuinely slower device still loses — the warm-up removes an
4220    /// artefact, it does not put a thumb on the scale.
4221    #[test]
4222    fn a_slow_device_still_loses_after_the_warmup() {
4223        let p = Probe::new();
4224        for _ in 0..4 {
4225            probe_record_into(&p, "matvec", None, true, ms(40.0));
4226        }
4227        for _ in 0..4 {
4228            probe_record_into(&p, "matvec", None, false, ms(2.0));
4229        }
4230        assert_eq!(p.state.load(Ordering::Relaxed), 2, "host wins on merit");
4231    }
4232}
4233
4234/// Scratch/weight lifetime for a synchronous image-pipeline stage. Declare
4235/// this before the stage model so the model drops before cache collection.
4236pub(crate) struct ImageStageGuard {
4237    #[cfg(target_os = "macos")]
4238    metal: Option<crate::gpu_metal::ImageStageGuard>,
4239    #[cfg(feature = "gpu")]
4240    wgpu: crate::gpu_wgpu::ImageStageGuard,
4241}
4242
4243pub(crate) fn image_stage_scope() -> ImageStageGuard {
4244    ImageStageGuard {
4245        #[cfg(target_os = "macos")]
4246        metal: if matches!(backend(), Backend::Metal) {
4247            Some(crate::gpu_metal::image_stage_scope())
4248        } else {
4249            None
4250        },
4251        #[cfg(feature = "gpu")]
4252        wgpu: crate::gpu_wgpu::image_stage_scope(),
4253    }
4254}
4255
4256impl ImageStageGuard {
4257    pub(crate) fn track_model(&mut self, uid: u64) {
4258        #[cfg(target_os = "macos")]
4259        if let Some(metal) = &mut self.metal {
4260            metal.track_model(uid);
4261        }
4262        #[cfg(feature = "gpu")]
4263        self.wgpu.track_model(uid);
4264        #[cfg(not(target_os = "macos"))]
4265        let _ = uid;
4266    }
4267}
4268
4269// ════════════════════════════════════════════════════════════════════
4270// Z-Image-Turbo device contract (WP0 scaffold, plan §2.1). APPEND-ONLY.
4271//
4272// Owner of the contract: the WP1 lead. The backends implement it in their
4273// own child modules — `gpu_wgpu/zimage.rs` (WP2) and `gpu_metal/zimage.rs`
4274// (WP3) — and never edit the parent files. New fields are added only as
4275// `Option<…>` with agreed semantics; existing fields never change meaning.
4276//
4277// Convention (the same as every `gpu::*` entry): `false` = "not handled",
4278// nothing observable was changed, and the caller runs the CPU path
4279// (`zimage::ZImageDit::step_cpu` etc.), which is the bit-level reference.
4280//
4281// Sequence order everywhere is diffusers' [img rows…, cap rows…], with
4282// padded lengths n_img_p = ceil32(n_img) and n_cap_p = ceil32(L).
4283// ════════════════════════════════════════════════════════════════════
4284
4285// ───────────────────── Qwen-Image-2.1 device contract ─────────────────────
4286
4287/// Geometry of the Qwen-Image-2.1 denoiser.
4288#[derive(Clone, Copy, Debug, PartialEq)]
4289pub struct Qi21Geom {
4290    pub hidden: usize,
4291    pub nh: usize,
4292    pub hd: usize,
4293    pub inter: usize,
4294    pub in_ch: usize,
4295    pub eps: f32,
4296}
4297
4298/// One block's weights: tensor indices of `to_q, to_k, to_v, to_out.0,
4299/// gate_layer, proj (up), out (down)` and the per-head q/k norm weights.
4300pub struct Qi21BlockRef<'a> {
4301    pub w: [usize; 7],
4302    pub norm_q: &'a [f32],
4303    pub norm_k: &'a [f32],
4304}
4305
4306/// Everything the device needs to build a program and run its prefix.
4307pub struct Qi21PrefillArgs<'a> {
4308    pub model: &'a Arc<CmfModel>,
4309    pub geom: Qi21Geom,
4310    pub blocks: &'a [Qi21BlockRef<'a>],
4311    /// `img_in` [hidden, 64] and `proj_out` [64, hidden], f32.
4312    pub img_in: &'a [f32],
4313    pub proj_out: &'a [f32],
4314    pub key: u64,
4315    /// Prefix rows after `txt_in` / `img_in`, `[lp, hidden]`.
4316    pub x: &'a [f32],
4317    pub lp: usize,
4318    /// RoPE cos/sin of the prefix rows and of the target rows, `[rows, 64]`.
4319    pub rope_p: (&'a [f32], &'a [f32]),
4320    pub rope_t: (&'a [f32], &'a [f32]),
4321    /// Keys each prefix row sees (a leading run; non-decreasing).
4322    pub vis: &'a [u32],
4323    /// The t = 0 modulation `[s1|g1|s2|g2]`.
4324    pub mods0: &'a [f32],
4325    /// Target rows.
4326    pub n: usize,
4327}
4328
4329/// Build the program `a.key` and run its prefix on the device.
4330#[allow(unused_variables)]
4331pub fn qi21_prefill(a: &Qi21PrefillArgs) -> bool {
4332    match backend() {
4333        #[cfg(target_os = "macos")]
4334        Backend::Metal => crate::gpu_metal::qi21::prefill(a),
4335        #[cfg(feature = "gpu")]
4336        Backend::Wgpu => crate::gpu_wgpu::qi21::prefill(a),
4337        #[allow(unreachable_patterns)]
4338        _ => false,
4339    }
4340}
4341
4342/// One denoiser call of program `key`: `xtok [n, 64]` → `out [n, 64]`.
4343#[allow(unused_variables)]
4344pub fn qi21_step(key: u64, xtok: &[f32], mods: &[f32], fs: &[f32], out: &mut [f32]) -> bool {
4345    match backend() {
4346        #[cfg(target_os = "macos")]
4347        Backend::Metal => crate::gpu_metal::qi21::step(key, xtok, mods, fs, out),
4348        #[cfg(feature = "gpu")]
4349        Backend::Wgpu => crate::gpu_wgpu::qi21::step(key, xtok, mods, fs, out),
4350        #[allow(unreachable_patterns)]
4351        _ => false,
4352    }
4353}
4354
4355/// Drop program `key`. The wgpu module is released directly (module-local
4356/// state only), so a release never brings a device up.
4357#[allow(unused_variables)]
4358pub fn qi21_release_key(key: u64) {
4359    #[cfg(target_os = "macos")]
4360    if matches!(backend(), Backend::Metal) {
4361        crate::gpu_metal::qi21::release_key(key);
4362    }
4363    #[cfg(feature = "gpu")]
4364    crate::gpu_wgpu::qi21::release_key(key);
4365}
4366
4367/// Drop the denoiser's device state (see `qi21_release_key`), and the
4368/// resident VAE decoder's weights.
4369pub fn qi21_release() {
4370    #[cfg(target_os = "macos")]
4371    if matches!(backend(), Backend::Metal) {
4372        crate::gpu_metal::qi21::release();
4373    }
4374    #[cfg(feature = "gpu")]
4375    {
4376        crate::gpu_wgpu::qi21::release();
4377        crate::gpu_wgpu::qi21_vae::release();
4378    }
4379}
4380
4381/// One conv of the Qwen-Image-2.1 VAE decoder: host f32 weights
4382/// `[co][ci][k][k]` and bias `[co]` (k = 1 or 3, stride 1, same padding).
4383#[derive(Clone, Copy)]
4384pub struct Qi21VaeConvRef<'a> {
4385    pub w: &'a [f32],
4386    pub b: &'a [f32],
4387    pub ci: usize,
4388    pub co: usize,
4389    pub k: usize,
4390}
4391
4392/// A residual block: `conv2(φ(conv1(φ(x)))) + x` (through the 1×1
4393/// shortcut when ci ≠ co), φ = RMS_norm over channels (x/‖x‖·√C·γ), SiLU.
4394pub struct Qi21VaeResRef<'a> {
4395    pub g1: &'a [f32],
4396    pub c1: Qi21VaeConvRef<'a>,
4397    pub g2: &'a [f32],
4398    pub c2: Qi21VaeConvRef<'a>,
4399    pub shortcut: Option<Qi21VaeConvRef<'a>>,
4400}
4401
4402/// An up block: its resnets, then (all but the last block) nearest-2× +
4403/// the 3×3 conv, plus the DupUp shortcut of the block's input (`in_dim` →
4404/// `out_dim` channels, temporal factor `ft`).
4405pub struct Qi21VaeUpRef<'a> {
4406    pub resnets: Vec<Qi21VaeResRef<'a>>,
4407    pub up: Option<(Qi21VaeConvRef<'a>, usize)>,
4408    pub in_dim: usize,
4409    pub out_dim: usize,
4410}
4411
4412/// The Qwen-Image-2.1 VAE decoder (`qwen_image21_vae.rs`) for a resident
4413/// device decode.
4414pub struct Qi21VaeDecodeArgs<'a> {
4415    /// Identity of the weights: a backend caches its planes by it.
4416    pub key: u64,
4417    pub post_quant: Qi21VaeConvRef<'a>,
4418    pub conv_in: Qi21VaeConvRef<'a>,
4419    pub mid_res: [Qi21VaeResRef<'a>; 2],
4420    /// The mid attention: RMS_norm γ, `to_qkv` (1×1, 3c outputs: q, k, v),
4421    /// `proj` (1×1); single head, residual.
4422    pub attn_gamma: &'a [f32],
4423    pub attn_qkv: Qi21VaeConvRef<'a>,
4424    pub attn_proj: Qi21VaeConvRef<'a>,
4425    pub ups: Vec<Qi21VaeUpRef<'a>>,
4426    pub norm_out: &'a [f32],
4427    pub conv_out: Qi21VaeConvRef<'a>,
4428}
4429
4430/// Resident VAE decode: `z` = raw latent planes `[z_dim, h·w]` → `out`
4431/// `[out_channels, 16h·16w]` (before the clamp). `false` = not handled,
4432/// `out` untouched: the per-conv path runs.
4433#[allow(unused_variables)]
4434pub fn qi21_vae_decode(a: &Qi21VaeDecodeArgs, z: &[f32], h: usize, w: usize, out: &mut [f32]) -> bool {
4435    match backend() {
4436        #[cfg(feature = "gpu")]
4437        Backend::Wgpu => crate::gpu_wgpu::qi21_vae::decode(a, z, h, w, out),
4438        #[allow(unreachable_patterns)]
4439        _ => false,
4440    }
4441}
4442
4443/// One Z-Image transformer block's device inputs (noise refiner, context
4444/// refiner or main layer — all share this shape). Weights are tensor
4445/// indices into `model.tensors` (diffusers names under `dit.`); the codec
4446/// is whatever the container holds (F16/Bf16/Q8Row/Q8_2f/Q4TiledP…), and a
4447/// backend that cannot expand a codec declines (returns `false`).
4448/// Norm vectors are f32 host slices that live as long as the caller's
4449/// `ZImageDit`; a backend may cache them by pointer (they do not change).
4450#[derive(Clone, Copy)]
4451pub struct ZBlockRef<'a> {
4452    /// `attention.to_q/to_k/to_v/to_out.0.weight`, each [hidden, hidden].
4453    pub wq: usize,
4454    pub wk: usize,
4455    pub wv: usize,
4456    pub wo: usize,
4457    /// `feed_forward.w1` (gate) / `w3` (up) [inter, hidden], `w2` (down)
4458    /// [hidden, inter]. FFN = w2(silu(w1·x) ⊙ w3·x).
4459    pub w1: usize,
4460    pub w3: usize,
4461    pub w2: usize,
4462    /// `attention_norm1` / `attention_norm2`, [hidden] (plain-w RMSNorm).
4463    pub norm1: &'a [f32],
4464    pub norm2: &'a [f32],
4465    /// `ffn_norm1` / `ffn_norm2`, [hidden].
4466    pub ffn_norm1: &'a [f32],
4467    pub ffn_norm2: &'a [f32],
4468    /// `attention.norm_q` / `norm_k`, [hd] (per-head RMSNorm before RoPE).
4469    pub norm_q: &'a [f32],
4470    pub norm_k: &'a [f32],
4471}
4472
4473/// Z-Image geometry. Turbo: hidden 3840, nh 30 (MHA, no GQA), hd 128,
4474/// inter 10240, eps 1e-5 (all RMSNorms incl. qk-norm), final_eps 1e-6
4475/// (the affine-free final LayerNorm), patch_dim 64 (2×2×16).
4476#[derive(Clone, Copy, Debug, PartialEq)]
4477pub struct ZGeom {
4478    pub hidden: usize,
4479    pub nh: usize,
4480    pub hd: usize,
4481    pub inter: usize,
4482    pub eps: f32,
4483    pub final_eps: f32,
4484    pub patch_dim: usize,
4485}
4486
4487/// Once per (prompt, resolution). The backend uploads/caches what it needs
4488/// keyed by `key`; weight planes are keyed by the MODEL (not by `key`) and
4489/// survive across prompts until `zimage_release`.
4490pub struct ZPrepareArgs<'a> {
4491    pub model: &'a Arc<CmfModel>,
4492    pub geom: ZGeom,
4493    /// Caller-chosen identity of this (prompt, resolution) state; every
4494    /// `ZStepArgs` of the same image carries the same key.
4495    pub key: u64,
4496    /// Image tokens (H/16 · W/16), padded count ceil32(n_img), caption
4497    /// padded count ceil32(L). S = n_img_p + n_cap_p.
4498    pub n_img: usize,
4499    pub n_img_p: usize,
4500    pub n_cap_p: usize,
4501    /// The patch grid (H/16, W/16); n_img = grid.0 · grid.1. Row-major
4502    /// token order `hp·grid.1 + wp`.
4503    pub grid: (usize, usize),
4504    /// [n_cap_p, hidden], ALREADY context-refined (host or device).
4505    pub cap: &'a [f32],
4506    /// Noise-refiner RoPE: [n_img_p · hd/2] cos, sin (complex-interleaved
4507    /// pairs, hd/2 angles per token).
4508    pub rope_img: (&'a [f32], &'a [f32]),
4509    /// Main-layer RoPE: [(n_img_p + n_cap_p) · hd/2], rows ordered [img, cap].
4510    pub rope_joint: (&'a [f32], &'a [f32]),
4511    /// `all_x_embedder.2-1.weight` [hidden, 64], `.bias` [hidden],
4512    /// `x_pad_token` [hidden] (replaces rows ≥ n_img after the embed).
4513    pub x_emb_w: &'a [f32],
4514    pub x_emb_b: &'a [f32],
4515    pub x_pad: &'a [f32],
4516    /// `all_final_layer.2-1.linear.weight` [64, hidden], `.bias` [64].
4517    pub final_w: &'a [f32],
4518    pub final_b: &'a [f32],
4519    /// 2 noise-refiner blocks (image rows only) and 30 main layers.
4520    pub noise_refiner: &'a [ZBlockRef<'a>],
4521    pub layers: &'a [ZBlockRef<'a>],
4522    /// OPTIONAL (backends may ignore): the modulation of EVERY step of this
4523    /// image, [steps][(2+30)·4·hidden] in the `ZStepArgs::mods` layout, and
4524    /// [steps][hidden] final scales, so a backend can upload them once per
4525    /// image and index them by `ZStepArgs::step`. `ZStepArgs::mods` is still
4526    /// always supplied and is authoritative.
4527    pub mods_all: Option<&'a [f32]>,
4528    pub final_scale_all: Option<&'a [f32]>,
4529    /// OPTIONAL (B2): the CFG negative item. When `Some`, the backend
4530    /// prepares ONE batch-2 program under `key` — item 0 is this prompt,
4531    /// item 1 the negative — and every `ZStepArgs` of that key must carry
4532    /// `out_neg`. A backend without batch 2 returns `false` (the caller
4533    /// then prepares the two items separately or runs the CPU path).
4534    pub neg: Option<ZNegArgs<'a>>,
4535}
4536
4537/// The negative (unconditional) item of a CFG pair: its own refined
4538/// caption, padded caption length and joint RoPE table (the image ids sit
4539/// at axis-0 position L_p+1, so both tables depend on the item's L_p).
4540pub struct ZNegArgs<'a> {
4541    /// [n_cap_p, hidden], context-refined.
4542    pub cap: &'a [f32],
4543    pub n_cap_p: usize,
4544    /// [n_img_p · hd/2] cos, sin (noise refiner) of this item.
4545    pub rope_img: (&'a [f32], &'a [f32]),
4546    /// [(n_img_p + n_cap_p) · hd/2] cos, sin, rows [img, cap].
4547    pub rope_joint: (&'a [f32], &'a [f32]),
4548}
4549
4550/// Once per denoising step.
4551pub struct ZStepArgs<'a> {
4552    /// The `ZPrepareArgs::key` this step belongs to. A key the backend has
4553    /// not prepared → `false`.
4554    pub key: u64,
4555    /// Step index into the schedule (0..steps); selects the row of
4556    /// `ZPrepareArgs::mods_all` when a backend uses it.
4557    pub step: usize,
4558    /// [n_img_p, 64] patchified latent, inner order (dy·2+dx)·16+c. Rows
4559    /// ≥ n_img are copies of the last row; the backend replaces them with
4560    /// `x_pad` after the embed.
4561    pub x_tok: &'a [f32],
4562    /// Per block (noise_refiner then layers) the RAW chunks
4563    /// [scale_msa, gate_msa, scale_mlp, gate_mlp] of Linear(temb) (no SiLU
4564    /// before it), [(2+30)·4·hidden]. The backend applies (1+s) and tanh(g).
4565    pub mods: &'a [f32],
4566    /// [hidden] = 1 + Linear(SiLU(temb)) — already includes the +1.
4567    pub final_scale: &'a [f32],
4568    /// [n_img, 64]: the model output v (before the pipeline's negation),
4569    /// image rows only, patchified order.
4570    pub out: &'a mut [f32],
4571    /// [n_img, 64]: the negative item's v — required (and only valid) for
4572    /// a key prepared with `ZPrepareArgs::neg`. Both items see `x_tok`.
4573    pub out_neg: Option<&'a mut [f32]>,
4574}
4575
4576/// Prepare the per-(prompt, resolution) device state. Backends: wgpu →
4577/// `gpu_wgpu::zimage::prepare` (WP2), Metal → `gpu_metal::zimage::prepare`
4578/// (WP3).
4579#[allow(unused_variables)]
4580pub fn zimage_prepare(a: &ZPrepareArgs) -> bool {
4581    match backend() {
4582        #[cfg(target_os = "macos")]
4583        Backend::Metal => crate::gpu_metal::zimage::prepare(a),
4584        #[cfg(feature = "gpu")]
4585        Backend::Wgpu => crate::gpu_wgpu::zimage::prepare(a),
4586        #[allow(unreachable_patterns)]
4587        _ => false,
4588    }
4589}
4590
4591/// One full DiT forward on the device: x_embed → pad rows → noise refiner
4592/// ×2 → concat [img, cap] → 30 layers → final LayerNorm·scale → Linear →
4593/// image rows into `a.out`.
4594#[allow(unused_variables)]
4595pub fn zimage_step(a: &mut ZStepArgs) -> bool {
4596    match backend() {
4597        #[cfg(target_os = "macos")]
4598        Backend::Metal => crate::gpu_metal::zimage::step(a),
4599        #[cfg(feature = "gpu")]
4600        Backend::Wgpu => crate::gpu_wgpu::zimage::step(a),
4601        #[allow(unreachable_patterns)]
4602        _ => false,
4603    }
4604}
4605
4606/// OPTIONAL (B2): build the backend's weight planes for the per-step blocks
4607/// and the context refiner ahead of `zimage_prepare`, so the caller can
4608/// overlap the upload with the (CPU) text encoder. `false` = not done;
4609/// `zimage_prepare` builds whatever is missing either way.
4610#[allow(unused_variables)]
4611pub fn zimage_preload(
4612    model: &Arc<CmfModel>,
4613    geom: &ZGeom,
4614    noise_refiner: &[ZBlockRef],
4615    layers: &[ZBlockRef],
4616    context_refiner: &[ZBlockRef],
4617) -> bool {
4618    match backend() {
4619        #[cfg(target_os = "macos")]
4620        Backend::Metal => {
4621            crate::gpu_metal::zimage::preload(model, geom, noise_refiner, layers, context_refiner)
4622        }
4623        #[cfg(feature = "gpu")]
4624        Backend::Wgpu => crate::gpu_wgpu::zimage::preload(model, geom, noise_refiner, layers, context_refiner),
4625        #[allow(unreachable_patterns)]
4626        _ => false,
4627    }
4628}
4629
4630/// Persist the driver's compiled pipelines after a Z-Image generation (the
4631/// chain's kernels are built at first use, after the context came up), so
4632/// the next process skips the compile. Best-effort, no-op off wgpu.
4633pub fn zimage_flush_pipelines() {
4634    #[cfg(feature = "gpu")]
4635    if matches!(backend(), Backend::Wgpu) {
4636        crate::gpu_wgpu::pipeline_cache_flush();
4637    }
4638}
4639
4640/// OPTIONAL (B2): bring the device up and compile the Z-Image kernels, on
4641/// a helper thread at the start of a generation (the context and the
4642/// compiles cost ~1 s cold, beside the host-side loading). `false` = no
4643/// device path here.
4644pub fn zimage_warmup() -> bool {
4645    match backend() {
4646        #[cfg(target_os = "macos")]
4647        Backend::Metal => crate::gpu_metal::zimage::warmup(),
4648        #[cfg(feature = "gpu")]
4649        Backend::Wgpu => crate::gpu_wgpu::zimage::warmup(),
4650        #[allow(unreachable_patterns)]
4651        _ => false,
4652    }
4653}
4654
4655/// OPTIONAL (B2): upload the resident VAE's weights and compile its
4656/// kernels ahead of `vae_decode_chain` (the caller runs it on a helper
4657/// thread while the DiT steps keep the device busy). `false` = not done.
4658#[allow(unused_variables)]
4659pub fn vae_prewarm(a: &crate::vae::VaeChainArgs) -> bool {
4660    match backend() {
4661        #[cfg(target_os = "macos")]
4662        Backend::Metal => crate::gpu_metal::zimage::vae_prewarm(a),
4663        #[cfg(feature = "gpu")]
4664        Backend::Wgpu => crate::gpu_wgpu::zimage::vae_prewarm(a),
4665        #[allow(unreachable_patterns)]
4666        _ => false,
4667    }
4668}
4669
4670/// Drop the Z-Image DiT device state (planes, prepared programs) but keep
4671/// the VAE chain (B2: the generator frees the DiT before decoding).
4672pub fn zimage_release_dit() {
4673    #[cfg(target_os = "macos")]
4674    crate::gpu_metal::zimage::release_dit();
4675    #[cfg(feature = "gpu")]
4676    crate::gpu_wgpu::zimage::release_dit();
4677}
4678
4679/// Drop every Z-Image device resource (planes, prepared states, VAE chain
4680/// buffers): stage change or process end. Calls each compiled backend's
4681/// release directly, without `backend()`, so it never brings a device up;
4682/// the child modules' `release` must touch module-local state only.
4683pub fn zimage_release() {
4684    #[cfg(target_os = "macos")]
4685    crate::gpu_metal::zimage::release();
4686    #[cfg(feature = "gpu")]
4687    crate::gpu_wgpu::zimage::release();
4688}
4689
4690/// Optional device context refiner: the same block math with scale = 0 and
4691/// gate = 1 (unmodulated): x += norm2(attn(norm1(x))); x += ffn_norm2(ffn(
4692/// ffn_norm1(x))). `cap` is [n_cap_p, hidden] in/out (the cap_embedder
4693/// output with pad rows already = cap_pad_token); `rope_cap` is
4694/// [n_cap_p · hd/2] cos, sin. `false` = untouched, run the CPU refiner.
4695#[allow(unused_variables)]
4696pub fn zimage_refine_caption(
4697    model: &Arc<CmfModel>,
4698    geom: &ZGeom,
4699    blocks: &[ZBlockRef],
4700    rope_cap: (&[f32], &[f32]),
4701    cap: &mut [f32],
4702) -> bool {
4703    match backend() {
4704        #[cfg(target_os = "macos")]
4705        Backend::Metal => {
4706            crate::gpu_metal::zimage::refine_caption(model, geom, blocks, rope_cap, cap)
4707        }
4708        #[cfg(feature = "gpu")]
4709        Backend::Wgpu => {
4710            crate::gpu_wgpu::zimage::refine_caption(model, geom, blocks, rope_cap, cap)
4711        }
4712        #[allow(unreachable_patterns)]
4713        _ => false,
4714    }
4715}
4716
4717/// Resident Flux-VAE decode (the whole decoder on the device, one latent
4718/// upload, one RGB readback). `a` comes from `VaeDecoder::chain_args()`.
4719/// `z` is [latent_channels, h, w] ALREADY de-normalised
4720/// (z/scaling_factor + shift_factor — the conv_in input); `out` is
4721/// [3, 8h, 8w], the raw decoder output (≈[-1, 1], before x/2+0.5).
4722#[allow(unused_variables)]
4723pub fn vae_decode_chain(
4724    a: &crate::vae::VaeChainArgs,
4725    z: &[f32],
4726    h: usize,
4727    w: usize,
4728    out: &mut [f32],
4729) -> bool {
4730    match backend() {
4731        #[cfg(target_os = "macos")]
4732        Backend::Metal => crate::gpu_metal::zimage::vae_decode_chain(a, z, h, w, out),
4733        #[cfg(feature = "gpu")]
4734        Backend::Wgpu => crate::gpu_wgpu::zimage::vae_decode_chain(a, z, h, w, out),
4735        #[allow(unreachable_patterns)]
4736        _ => false,
4737    }
4738}