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