Skip to main content

cortiq_engine/
pipeline.rs

1//! Full inference pipeline: tokenize → embed → layers → lm_head → sample → decode.
2//!
3//! Prefill/decode contract: every token is forwarded exactly once and
4//! enters the KV cache exactly once. Logits for the next token are
5//! computed from the hidden state of the LAST forwarded token — the
6//! decode loop forwards the freshly sampled token, never re-embeds the
7//! prompt tail (v1 duplicated the last prompt token in the cache).
8
9use crate::attention::{self, QwenAttnCfg};
10use crate::inference;
11
12/// MiMo-V2 multi-token prediction (draft stack + speculative round). A
13/// child module so it runs on the pipeline's own helpers.
14#[path = "mimo_mtp.rs"]
15pub mod mimo_mtp;
16use crate::kv_cache::KvCache;
17use crate::linear_core::{
18    GdnCfg, GdnWeights, ShortConvCfg, ShortConvWeights, VmfPhaseCfg, VmfPhaseWeights, gdn_forward,
19    gdn_pair, short_conv_forward, short_conv_forward_batch, short_conv_pair, vmf_phase_forward,
20    vmf_phase_pair,
21};
22use crate::pool::Pool;
23use crate::qtensor::QTensor;
24use crate::sampler::{self, SamplerConfig, SamplerScratch, SplitMix64};
25use crate::tokenizer::Tokenizer;
26use cortiq_core::mask::TaskMask;
27use cortiq_core::types::NormStyle;
28
29pub static GLOBAL_USE_GPU: std::sync::atomic::AtomicBool =
30    std::sync::atomic::AtomicBool::new(false);
31
32/// Reusable per-pipeline forward scratch: the four norm outputs the
33/// decode paths recompute every layer (single: n1/p1; pair: all four).
34/// Plain buffers, resized once — steady-state decode reuses them.
35struct ForwardScratch {
36    n1: Vec<f32>,
37    n2: Vec<f32>,
38    p1: Vec<f32>,
39    p2: Vec<f32>,
40}
41
42impl ForwardScratch {
43    fn new(hidden: usize) -> Self {
44        Self {
45            n1: vec![0.0; hidden],
46            n2: vec![0.0; hidden],
47            p1: vec![0.0; hidden],
48            p2: vec![0.0; hidden],
49        }
50    }
51}
52
53/// Complete inference pipeline state.
54pub struct Pipeline {
55    /// In-process layer split across local GPUs: (device, first layer,
56    /// last layer) per segment, in execution order. `None` = one device.
57    /// Arc so cloning the plan out of `&mut self` does not fight the
58    /// borrow checker on the hot path.
59    gpu_plan: Option<std::sync::Arc<Vec<(usize, usize, usize)>>>,
60    /// Arc: the server shares one tokenizer handle across request
61    /// handlers without borrowing a pipeline slot.
62    pub tokenizer: std::sync::Arc<Tokenizer>,
63    pub kv_cache: KvCache,
64    pub sampler_config: SamplerConfig,
65    pub weights: PipelineWeights,
66    pub hidden_size: usize,
67    pub intermediate_size: usize,
68    pub num_heads: usize,
69    pub num_kv_heads: usize,
70    pub head_dim: usize,
71    /// Total virtual layers (num_layers × num_loops for looped models).
72    pub num_layers: usize,
73    /// Physical layers in weights.layers (≤ num_layers for looped models).
74    pub physical_layers: usize,
75    /// Looped Transformer: apply final norm after each loop iteration.
76    pub loop_final_norm: bool,
77    pub vocab_size: usize,
78    pub rms_eps: f64,
79    pub rope_base: f32,
80    pub norm_style: NormStyle,
81    /// RoPE dims actually rotated (≤ head_dim; Qwen3.5 uses head_dim/4).
82    pub rotary_dim: usize,
83    /// Optional Q-head count override for each attention layer (Laguna).
84    pub attention_heads_per_layer: Option<Vec<usize>>,
85    /// Optional KV-head count of each PHYSICAL attention layer (MiMo-V2:
86    /// 4 on full-attention layers, 8 on sliding ones). Set only through
87    /// [`Pipeline::set_attn_geometry`], which also reshapes the layer
88    /// caches. None = every layer has `num_kv_heads`.
89    pub kv_heads_per_layer: Option<Vec<usize>>,
90    /// Width of each V head when it is narrower than `head_dim` (MiMo-V2:
91    /// 128 against 192). V is zero-padded to `head_dim` inside the cache
92    /// and the attention output is compacted back to nh·v_head_dim before
93    /// o_proj (see `QwenAttnCfg::v_head_dim`). None = `head_dim`.
94    pub v_head_dim: Option<usize>,
95    /// `CMF_LAYER_DUMP=<dir>` (read once at construction; tests set it
96    /// directly): the hidden state after every layer, for every position,
97    /// as raw little-endian f32 files `p{pos:06}_l{li:02}.f32` of
98    /// `hidden_size` floats each. Written by the CPU layer walks — the
99    /// batched prefill (`prefill_batch_span`) and the single-token forward
100    /// (`forward_layers_span`) — so both prompt ingest and decode can be
101    /// diffed layer by layer against an external oracle (tools/mimo_ref).
102    /// Layers a device graph runs (wgpu token/batch graph, Metal chunk or
103    /// block graphs) and the own-stack families (DeepSeek-V4/V4.1,
104    /// Qwen3.8-Flash-Next, Gemma-3n) are not dumped: run with CMF_GPU=0
105    /// for a complete set. `li` is the virtual layer index. Final logits:
106    /// `CMF_LOGIT_DUMP=<file>` (hidden + logits of the first decode step).
107    pub layer_dump: Option<std::path::PathBuf>,
108    /// GPU-graph declines already logged for this pipeline, as (graph
109    /// site, reason) — one line each, see `graph_attn_decline_reason`.
110    graph_declines: std::cell::RefCell<Vec<(&'static str, &'static str)>>,
111    /// MiMo-V2 expert placement (prefix / dynamic bank / hybrid), decided
112    /// on the first forward — see `crate::mimo_moe`.
113    pub(crate) mimo_moe: crate::mimo_moe::Slot,
114    /// Linear-core geometry (present when the model has linear layers).
115    pub vmf_cfg: Option<VmfPhaseCfg>,
116    /// GatedDeltaNet geometry (faithful vendor operator).
117    pub gdn_cfg: Option<GdnCfg>,
118    /// MiniCPM-class logit scale (tied lm_head → cannot fold into weights).
119    pub logit_multiplier: Option<f32>,
120    /// Cooperative cancel: set from any thread (FFI `cortiq_cancel`,
121    /// a dropped server connection); the generate loop checks it at
122    /// every prefill chunk and decode step and finishes with
123    /// `finish_reason: "cancelled"`. Auto-cleared when honoured.
124    pub cancel: std::sync::Arc<std::sync::atomic::AtomicBool>,
125    /// A GPU graph failure is distinct from a user/request cancellation.
126    /// Graph code sets this before raising the cooperative cancel flag so the
127    /// generation API can return an error instead of reporting a successful
128    /// `finish_reason: cancelled` result.
129    graph_failed: std::sync::atomic::AtomicBool,
130    /// Token ids currently materialized in the KV cache (the forwarded
131    /// prompt + all generated tokens except the last, which is sampled
132    /// but not yet forwarded). Lets the next generate call prefill only
133    /// the suffix when a chat app resends the whole history.
134    pub kv_history: Vec<u32>,
135    /// Owner tag of `kv_history`: the resident device graph held that
136    /// sequence (true) or the host did. Reuse continues only on the owner.
137    pub kv_history_device: bool,
138    /// KDA geometry (Kimi Linear / Kimi-K3) — shared by every Kda layer.
139    pub kda_cfg: Option<crate::linear_core::KdaCfg>,
140    /// Gemma-3n stack (AltUp/LAuReL/PLE/KV-sharing): its own forward —
141    /// weights.layers stays empty, the KV caches are the shared ones.
142    pub g3n: Option<Box<(crate::g3n::G3nGlobals, Vec<crate::g3n::G3nLayer>)>>,
143    /// DeepSeek-V4 runs its own stack too: its hidden state is `hc_mult`
144    /// copies of a vector, so no loop written for a single residual
145    /// stream can carry it.
146    pub dsv4: Option<
147        Box<(
148            crate::dsv4::Dsv4Globals,
149            Vec<crate::dsv4::Dsv4Layer>,
150            crate::dsv4::Dsv4Cfg,
151            crate::dsv4::Dsv4State,
152        )>,
153    >,
154    /// DeepSeek-V4.1 owns the shared CED/CSA2 attention state, raw Engram
155    /// lookup and four-stream mHC handoff. It cannot use the V4 cache
156    /// layout, so it has a dedicated executor and state tuple.
157    pub dsv41: Option<
158        Box<(
159            crate::dsv41::Dsv41Globals,
160            Vec<crate::dsv41::Dsv41Layer>,
161            crate::dsv41::Dsv41Cfg,
162            crate::dsv41::Dsv41State,
163        )>,
164    >,
165    /// Optional V4.1 vision tower. Text-only files leave this unset.
166    pub dsv41_vision: Option<crate::dsv41_vision::VisionModel>,
167    /// Prepared image rows consumed by the next V4.1 prefill.
168    dsv41_prefill: Option<(Vec<Option<Vec<f32>>>, Vec<bool>)>,
169    /// Qwen3.8-Flash-Next owns four residual streams plus QSA/PLE state;
170    /// the generic single-residual layer loop cannot represent it.
171    pub qwen4_exp: Option<
172        Box<(
173            crate::qwen4_exp::Globals,
174            Vec<crate::qwen4_exp::Layer>,
175            crate::qwen4_exp::Cfg,
176            crate::qwen4_exp::State,
177        )>,
178    >,
179    /// DeepSeek-V4's own speculation stack: three draft modules, each a full
180    /// layer, plus a confidence head on the last. Empty when the file has
181    /// none, which is the only signal the decode path needs.
182    pub dsv4_mtp: Vec<crate::dsv4::Dsv4Mtp>,
183    /// The draft's per-sequence state (KV rings, captured trunk hidden).
184    pub dspark: Option<crate::dsv4::DsparkState>,
185    /// Drafts awaiting their verdict: (position, proposals, still matching,
186    /// accepted so far).
187    pub dspark_pending: Vec<(usize, Vec<u32>, bool, usize)>,
188    /// Accepted prefix length of every graded draft.
189    pub dspark_hist: Vec<usize>,
190    /// The real tokens the drafts were graded against — a degenerate,
191    /// repeating output would make any acceptance number meaningless, and
192    /// the cheapest guard against believing one is to count them.
193    pub dspark_real: Vec<u32>,
194    /// The trunk's expert picks for the last few tokens, per layer. The
195    /// union over a window of them is what a batched verify would have to
196    /// read, and the ratio to the pick count is all it could save.
197    pub dspark_trunk_picks: Vec<Vec<(usize, Vec<usize>)>>,
198    /// (unique, total) expert picks per draft, trunk side and draft side.
199    pub dspark_exp: Vec<(usize, usize, usize, usize)>,
200    /// Wall time spent in the deliberately out-of-core draft. Kept separate
201    /// from trunk decode so block batching can be judged without conflating
202    /// it with GPU chain variance.
203    pub dspark_draft_ns: u128,
204    /// LFM2 short-convolution geometry (present when the model has
205    /// `ShortConv` mixer layers).
206    pub short_conv_cfg: Option<ShortConvCfg>,
207    /// Multi-token-prediction head (None = absent).
208    pub mtp: Option<MtpModule>,
209    /// MiMo-V2's draft stack (three chained MTP layers from the
210    /// `<stem>.mtp.cmf` sidecar); None = absent. Speculative greedy decode
211    /// uses it unless `CMF_MTP=0` / `CMF_MIMO_MTP=0`.
212    pub mimo_mtp: Option<mimo_mtp::MimoMtp>,
213    /// Set while the MiMo speculative verify runs `prefill_batch`: its MoE
214    /// layers take `moe_ffn_rows_exact` (each row bit-identical to decode).
215    verify_exact_moe: bool,
216    /// Speculative decode via MTP (greedy only; `CMF_MTP=0` disables).
217    pub speculative: bool,
218    /// Keep generating past end-of-sequence ids (the llama-bench contract
219    /// for a timed run). A loop flag, deliberately NOT a sampler
220    /// suppression: suppressed ids count as a penalty and switch the
221    /// speculative round and the greedy burst off, so a benchmark that
222    /// suppressed EOS never measured either.
223    pub ignore_eos: bool,
224    /// Draft-head shortlist guard: tokens left during which the draft
225    /// uses the FULL head because a recently committed id lay past the
226    /// `CMF_DRAFT_VOCAB` cut (Cyrillic and CJK ids sit above 131072 in
227    /// Qwen's table, so a prefix shortlist would draft nothing usable
228    /// there — measured on Russian prose: 2.9 → 1.6 accepted a round).
229    pub draft_full_streak: u32,
230    /// Adaptive draft depth for the speculative round (None until the
231    /// first round): grows while nearly every draft is accepted, shrinks
232    /// when fewer than half are. The verify's cost climbs with the rows on
233    /// a discrete card (RTX PRO 4000: 52 ms at 2 rows, 74 at 5, 80 at 6),
234    /// so prose wants k≈3 and code or the repetitive bench k≈5 — measured
235    /// 33.6 vs 27.6 tok/s on an essay at k=3 vs 5, 45.6 vs 38 on code.
236    /// `CMF_GRAPH_SPEC_K` pins it.
237    pub spec_k_adapt: Option<usize>,
238    /// EWMA of the accepted fraction that drives `spec_k_adapt`.
239    pub spec_acc_ewma: f32,
240    rng: SplitMix64,
241    sampler_scratch: SamplerScratch,
242    /// Speculative SAMPLING state (graph_spec_step, temperature > 0): the
243    /// correction token a rejected draft produced — committed by the loop
244    /// top in place of a fresh draw — and the per-round draft
245    /// distributions / target scratch, reused so a round allocates
246    /// nothing at the vocab size.
247    spec_forced: Option<u32>,
248    spec_q: Vec<Vec<f32>>,
249    spec_p: Vec<f32>,
250    spec_res: Vec<f32>,
251    /// The same three for the sparse chain (top-k configs).
252    spec_qs: Vec<sampler::Sparse>,
253    spec_ps: sampler::Sparse,
254    spec_ress: sampler::Sparse,
255    /// Which arm the MTP draft block runs on this generation: Some(true)
256    /// = the whole-token graph (device attention, one submit a step),
257    /// Some(false) = the per-op path; None = not decided yet. Decided
258    /// on the first draft and held, because the two arms keep the MTP
259    /// KV in different places (device mirror vs the CPU cache) and a
260    /// mid-run switch would read the wrong one.
261    mtp_graph_mode: Option<bool>,
262    /// The Metal verify graph of the round in flight, between its sync
263    /// (logits read) and the commit that replays the accepted prefix.
264    #[cfg(target_os = "macos")]
265    metal_verify: Option<MetalVerifyPending>,
266    /// Precomputed RoPE inverse frequencies [head_dim/2]. Arc: the
267    /// forward path clones a handle to escape the &mut self borrow —
268    /// cloning the table itself was a per-forward allocation.
269    pub(crate) inv_freq: std::sync::Arc<Vec<f32>>,
270    /// Reusable norm buffers for the decode hot path (roadmap §3 P0:
271    /// steady-state forward should not heap-allocate). Disjoint field
272    /// from `weights`/`kv_cache`, so split borrows keep working.
273    ws: ForwardScratch,
274    /// Persistent worker pool (None = serial; see CMF_THREADS).
275    pool: Option<std::sync::Arc<Pool>>,
276    // ── Dynamic per-token skill routing (spec §9, claim 14/16) ──
277    /// Source model, retained so a skill switch can re-resolve the
278    /// touched layers' FFN tensors (Mapped = mmap pointers, cheap).
279    pub(crate) model: Option<std::sync::Arc<cortiq_core::CmfModel>>,
280    /// Masks present → weights are dequantized f32 (rebuild path).
281    pub(crate) dyn_force_f32: bool,
282    /// Per-skill FFN layers actually replaced (derived from tensors, not
283    /// the meta `layers` field — ru2 replaces down_proj in 0..23 while
284    /// its meta says [20..23]). None = skill touches non-FFN tensors →
285    /// ineligible for cheap dynamic switching (honest refusal).
286    pub(crate) dyn_skill_layers: Vec<Option<Vec<usize>>>,
287    /// Currently overlaid skill (index into model.header.skills); None =
288    /// backbone. Set at load time to the statically-overlaid skill so
289    /// `set_active_skill(None)` correctly reverts it (else a static
290    /// skill would silently persist — the union-diff assumes dyn_active
291    /// always mirrors the live overlay). Switched by `set_active_skill`.
292    pub(crate) dyn_active: Option<usize>,
293    /// Pipeline was loaded with a soft blend (materialized working
294    /// tensors, not a single skill index) → dynamic routing refuses:
295    /// there is no single index to revert the blend from.
296    pub(crate) dyn_blend_loaded: bool,
297    /// Layer whose post-residual hidden feeds the router φ (shared by
298    /// swarm skills). None = φ capture off.
299    pub(crate) dyn_phi_layer: Option<usize>,
300    /// EMA of φ at `dyn_phi_layer` over the decode window (on-policy).
301    dyn_phi_ema: Vec<f32>,
302    dyn_phi_seen: usize,
303    /// Hysteresis router driving per-token skill switches during decode
304    /// (None = static/no dynamic routing). Taken out during generation.
305    pub dyn_router: Option<crate::swarm::DynRouter>,
306    /// O(1) Nyström attention setting (CLI/env/header-hint resolved by
307    /// the caller; None = plain cache attention everywhere).
308    o1_cfg: Option<crate::nystrom::O1Cfg>,
309    /// Bumped once per collecting→sealed transition — the GPU state mirror
310    /// re-uploads when it sees a new epoch (each fresh sealed state).
311    o1_epoch: u64,
312    /// Per-layer o1 flags derived from `o1_cfg` (Full layers only).
313    o1_flags: Vec<bool>,
314    /// Emit a structured per-token trace (B4 telemetry channel). Off by
315    /// default — the runtime is silent unless observation is requested.
316    trace: bool,
317    /// Confidence-calibration temperature (B1): reported probability is
318    /// softmax(logits / calib_temp). 1.0 = raw. Set from header.calibration.
319    calib_temp: f32,
320    /// Process-unique id keying this pipeline's device KV mirrors.
321    #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
322    graph_kv_id: u64,
323    /// Decode asks the token graph to also run final-norm + lm_head on
324    /// the device (drops the separate per-op lm_head round trip).
325    #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
326    graph_want_logits: bool,
327    /// NLL quality gates require the graph's fused head rather than silently
328    /// accepting a CPU head fallback. Generation keeps the historical
329    /// best-effort `graph_want_logits` behavior.
330    #[cfg_attr(not(target_os = "macos"), allow(dead_code))]
331    graph_head_required: bool,
332    /// Logits the graph produced for the token just forwarded (taken by
333    /// the decode loop; None = compute on the CPU path).
334    graph_logits: Option<Vec<f32>>,
335    /// Packed Embryo graph model, built lazily on the first eligible
336    /// resident call so other architectures pay no packing cost.
337    embryo_graph: Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>>,
338    /// This pipeline's token graph refused for a STRUCTURAL reason (its
339    /// own weights / layer kinds): not retried. Per pipeline, never
340    /// process-wide: a skill lane whose pack does not build must not flip
341    /// the backbone slot mid-sequence onto the host path (and a new lane
342    /// must not clear the backbone's verdict) — R4/NF-2.
343    graph_refused: std::sync::atomic::AtomicBool,
344    /// Token embeddings are multiplied by this at input (Gemma: √hidden).
345    pub embed_multiplier: f32,
346    /// Attention score scale (1/√head_dim unless the arch overrides —
347    /// Gemma's query_pre_attn_scalar).
348    pub attn_scale: f32,
349    /// Sliding-window attention: (window, every-Nth-layer-is-global
350    /// pattern) — Gemma-3.
351    pub swa: Option<(usize, usize)>,
352    /// Explicit local/global schedule for architectures that cannot be
353    /// represented by Gemma's every-Nth-global convention.
354    pub sliding_layers: Option<Vec<bool>>,
355    /// Natively bounded anchor record (`arch.anchor_core`): the file's
356    /// operator, installed once at load. `Some` = the model is
357    /// bounded-native — `--o1`/`CMF_O1*` are refused, prefix reuse and
358    /// the penalty window are bounded, and no anchor layer stores
359    /// anything per position.
360    pub anchor_core: Option<cortiq_core::AnchorCoreConfig>,
361    /// The `[W][rd/2]` relative-rotation table every bounded layer shares
362    /// (built from `inv_freq` once the RoPE setup is final).
363    bounded_rope: Option<std::sync::Arc<crate::bounded::BoundedRope>>,
364    /// Bounded prefix-reuse key (length + rolling hash + tail) of a
365    /// bounded-native model; `kv_history` stays empty there so nothing
366    /// in the pipeline grows with the dialogue.
367    pub kv_prefix: KvPrefix,
368    /// Prompt positions the last `generate*` call actually forwarded
369    /// (prefix reuse subtracts what the cache already held).
370    pub last_prefill_tokens: usize,
371    /// RoPE table of the sliding (local) layers, when they use their
372    /// own base frequency (Gemma-3: 10k local vs 1M global).
373    pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
374    pub rotary_dim_local: Option<usize>,
375    pub rope_scale: f32,
376    pub rope_scale_local: f32,
377    /// Gemma-4: global layers run their own geometry — (head_dim,
378    /// num_kv_heads); sliding layers keep the base fields.
379    pub global_attn: Option<(usize, usize)>,
380    /// Gemma-4: the global layers' proportional RoPE table (len
381    /// global_head_dim/2, zero-padded tail = identity rotation).
382    pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
383    /// Scale-less RMS normalization of V heads before caching (Gemma-4).
384    pub attn_v_norm: bool,
385    /// HunYuan dense: per-head q/k norm runs after RoPE (see the arch flag).
386    pub qk_norm_after_rope: bool,
387    /// The per-head `self_attn.g_proj` output gate uses sigmoid
388    /// (Spark-X2.5) instead of softplus (Laguna).
389    pub proj_gate_sigmoid: bool,
390    /// Final-logit soft-capping C: logits = C·tanh(logits/C) (Gemma-4).
391    pub final_softcap: Option<f32>,
392    /// Cortiq Embryo hierarchical head: cluster matrix [C, hidden]. The
393    /// flat logits h·Eᵀ are turned into the two-level log-probabilities
394    /// log softmax_c(h·Cᵀ)[c(v)] + log softmax_{s∈c(v)}(h·E_c(v)ᵀ)[v].
395    pub head_clusters: Option<std::sync::Arc<Vec<f32>>>,
396    /// Gemma-2 attention-logit soft-capping (0.0 = off).
397    pub attn_softcap: f32,
398    /// Compute per-token confidence (a full-vocab softmax each
399    /// token). On by default; `bench --core` turns it off to match
400    /// llama-bench's core timing.
401    confidence_on: bool,
402    /// Test-only one-shot forward failure, scoped to this pipeline so
403    /// parallel scoring tests cannot consume one another's injection.
404    #[cfg(test)]
405    nll_test_fail_at: Option<usize>,
406    /// Test-only route override; avoids mutating the process-wide
407    /// `CMF_PREFILL` environment variable while forcing the serial path.
408    #[cfg(test)]
409    nll_test_force_serial: bool,
410}
411
412#[cfg(target_os = "macos")]
413impl Drop for Pipeline {
414    fn drop(&mut self) {
415        // the async replay writes into `kv_cache` Vecs about to be freed
416        let _ = crate::gpu_metal::wait_replay();
417        crate::gpu::kv_mirror_drop(self.graph_kv_id);
418    }
419}
420
421#[cfg(not(target_os = "macos"))]
422impl Drop for Pipeline {
423    fn drop(&mut self) {
424        // The wgpu resident Embryo graph owns recurrent/KV buffers keyed by
425        // this pipeline's sequence id.  Release that sequence image when a
426        // pooled pipeline is dropped; model weights stay cached for reuse.
427        crate::gpu::graph_kv_reset(self.graph_kv_id);
428    }
429}
430
431/// Model weights. Matrices are `QTensor` (owned f32 for small models
432/// and tests — bit-identical to the historical paths — or quantized
433/// bytes zero-copy from the CMF mmap for big models). 1-D norms are
434/// always small and stay f32.
435pub struct PipelineWeights {
436    /// Embedding table: [vocab_size, hidden_size]
437    pub embed_tokens: QTensor,
438    /// Per-layer weights
439    pub layers: Vec<LayerWeights>,
440    /// LM head: [vocab_size, hidden_size]
441    pub lm_head: QTensor,
442    /// Final norm: [hidden_size]
443    pub final_norm: Vec<f32>,
444}
445
446/// One transformer layer: shared norms + MLP, attention by kind.
447pub struct LayerWeights {
448    pub input_norm: Vec<f32>,
449    /// The pre-FFN norm (`post_attention_layernorm` classically;
450    /// `pre_feedforward_layernorm` on Gemma-2/3 sandwich layers).
451    pub post_norm: Vec<f32>,
452    /// Gemma-2/3 sandwich: norm applied to the ATTENTION OUTPUT before
453    /// its residual add (`post_attention_layernorm` there).
454    pub attn_out_norm: Option<Vec<f32>>,
455    /// Gemma-4: the whole layer output is multiplied by this scalar.
456    pub layer_scale: Option<f32>,
457    /// Gemma-2/3 sandwich: norm applied to the FFN OUTPUT before its
458    /// residual add (`post_feedforward_layernorm`).
459    pub ffn_out_norm: Option<Vec<f32>>,
460    pub ffn: FfnKind,
461    pub attn: AttnKind,
462}
463
464/// FFN gate activation: SiLU (SwiGLU family) or tanh-GELU (Gemma's
465/// GeGLU). A property of the model, carried on every FFN triple.
466#[derive(Clone, Copy, PartialEq, Debug, Default)]
467pub enum Act {
468    #[default]
469    Silu,
470    GeluTanh,
471    /// Exact erf GELU (HF `hidden_act = "gelu"`: Spark-X2.5).
472    Gelu,
473    /// Kimi-K3 SituAndMul: BOTH halves transform —
474    /// a = β·tanh(g/β)·σ(g), up' = linβ·tanh(u/linβ) (linβ>0), out = a·up'.
475    Situ {
476        beta: f32,
477        linear_beta: f32,
478    },
479}
480
481impl Act {
482    pub fn from_arch(name: &str) -> Self {
483        if name == "gelu_tanh" {
484            Self::GeluTanh
485        } else if name == "gelu" {
486            Self::Gelu
487        } else {
488            Self::Silu
489        }
490    }
491
492    /// Arch-driven constructor (activation name + situ betas).
493    pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
494        match arch.hidden_act.as_str() {
495            "situ" => Self::Situ {
496                beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
497                linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
498            },
499            other => Self::from_arch(other),
500        }
501    }
502
503    #[inline]
504    pub fn apply(self, x: f32) -> f32 {
505        match self {
506            Self::Silu => inference::silu(x),
507            Self::GeluTanh => inference::gelu_tanh(x),
508            Self::Gelu => inference::gelu_erf(x),
509            Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
510        }
511    }
512
513    /// Gated combine — the FFN contract. Situ transforms the UP half
514    /// too, so callers must use this instead of apply(g)·u.
515    #[inline]
516    pub fn combine(self, g: f32, u: f32) -> f32 {
517        match self {
518            Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
519                self.apply(g) * (linear_beta * (u / linear_beta).tanh())
520            }
521            _ => self.apply(g) * u,
522        }
523    }
524
525    /// The wgpu graphs' arm for this activation; None = no device kernel
526    /// computes it, and the graph builder must refuse the layer (the dense
527    /// graph FFN once computed SiLU for whatever the model asked for).
528    pub fn graph_act(self) -> Option<crate::gpu::GraphAct> {
529        match self {
530            Self::Silu => Some(crate::gpu::GraphAct::Silu),
531            Self::Gelu => Some(crate::gpu::GraphAct::GeluErf),
532            Self::GeluTanh | Self::Situ { .. } => None,
533        }
534    }
535}
536
537/// Dense gated triple — the FFN of a dense layer or of one expert.
538pub struct DenseFfn {
539    pub gate_proj: QTensor,
540    pub up_proj: QTensor,
541    pub down_proj: QTensor,
542    /// Gate activation (SiLU default; Gemma: tanh-GELU).
543    pub act: Act,
544    /// `down_proj` stored transposed (`[inter, hidden]`), when the file
545    /// carries it. Only the per-token sparse path reads it: a neuron's
546    /// down weights are a contiguous ROW there, so the token's chosen
547    /// neurons are the only bytes touched. `None` = the ordinary layout,
548    /// and the sparse path stays off.
549    pub down_t: Option<QTensor>,
550    /// Task tubes (spec: defragged task-conditional width). The three
551    /// matrices above are the CORE — the neurons every task computes;
552    /// each tube is an independently quantized slice of the SAME layer
553    /// holding the neurons only some tasks need. A tube is a normal
554    /// tensor triple, so every kernel runs it unchanged, and the bytes
555    /// of an inactive tube are never read. Empty = ordinary dense FFN.
556    pub segs: Vec<FfnSeg>,
557}
558
559/// One task tube: a contiguous slice of a layer's FFN neurons, stored
560/// as its own `[w, hidden]` / `[hidden, w]` triple. `start` is the
561/// neuron's index in the layer's FULL space (core first, then tubes in
562/// order) — the bit a task mask sets to switch this tube on.
563pub struct FfnSeg {
564    pub gate: QTensor,
565    pub up: QTensor,
566    pub down: QTensor,
567    pub start: usize,
568    pub width: usize,
569}
570
571/// FFN operator of a layer, decided by tensor presence at load time
572/// (router `mlp.gate.weight` in the directory = MoE layer).
573pub enum FfnKind {
574    Dense(DenseFfn),
575    /// Mixture-of-Experts (Qwen2-MoE / Qwen3-MoE): softmax over ALL
576    /// expert logits → top-k, optional renorm; experts stay quantized
577    /// in mmap — only the selected ones are touched per token.
578    Moe(MoeFfn),
579    /// Gemma-4 MoE: a dense MLP branch AND a routed-expert branch in
580    /// the SAME layer, each with its own norm sandwich. The dense
581    /// branch reads the pre-FFN-normed input; the expert branch (and
582    /// the router) read the RAW residual through `pre_norm_2`:
583    ///   d = post_norm_1(dense(x̂));  m = post_norm_2(Σwₑ·FFNₑ(pre_norm_2(h)))
584    ///   ffn_out = d + m   (the caller's ffn_out_norm + residual follow)
585    DenseMoe(Box<DenseMoeFfn>),
586}
587
588/// Gemma-4 dual-branch FFN (see `FfnKind::DenseMoe`).
589pub struct DenseMoeFfn {
590    pub dense: DenseFfn,
591    pub moe: MoeFfn,
592    /// post_feedforward_layernorm_1 — dense-branch output norm.
593    pub post_norm_1: Vec<f32>,
594    /// pre_feedforward_layernorm_2 — expert-branch input norm (applied
595    /// to the RAW residual, not the pre-FFN-normed activation).
596    pub pre_norm_2: Vec<f32>,
597    /// post_feedforward_layernorm_2 — expert-branch output norm.
598    pub post_norm_2: Vec<f32>,
599}
600
601pub struct MoeFfn {
602    /// Router `mlp.gate.weight` [num_experts, hidden].
603    pub router: QTensor,
604    pub experts: Vec<DenseFfn>,
605    pub top_k: usize,
606    pub norm_topk_prob: bool,
607    /// Router scores per-expert with a sigmoid (LFM2-MoE / DeepSeek-V3
608    /// `noaux_tc`) instead of a softmax over all experts (Qwen).
609    pub router_sigmoid: bool,
610    /// Per-expert selection bias `mlp.expert_bias` [num_experts]
611    /// (LFM2-MoE): added to the sigmoid scores for the top-k CHOICE only;
612    /// the gathered weights use the unbiased scores. None = no bias.
613    pub expert_bias: Option<Vec<f32>>,
614    /// Top-k weights are multiplied by this after the optional renorm
615    /// (LFM2-MoE `routed_scaling_factor`; 1.0 = off).
616    pub routed_scaling: f32,
617    /// Adaptive routing (CMF_MOE_TAU, opt-in): keep the smallest
618    /// prefix of the top-k whose renormalized mass reaches τ —
619    /// confident tokens touch 1–2 experts, flat ones keep all k.
620    /// MoE decode is memory-bound, so skipped experts are skipped
621    /// weight traffic. None = classic fixed top-k (bit-identical).
622    pub route_tau: Option<f32>,
623    /// Always-on shared expert. Qwen2-MoE carries an additional sigmoid
624    /// gate; Laguna adds the shared expert unconditionally (`None`).
625    pub shared: Option<(DenseFfn, Option<QTensor>)>,
626    /// Expert-selection counters (truncated Fisher B-field of claim 12:
627    /// routing frequency during calibration). Filled by every forward,
628    /// read by the CLI via CMF_MOE_STATS. RefCell: decode is single-threaded.
629    pub stats: std::cell::RefCell<Vec<u64>>,
630    /// Per-CHANNEL sum of squares of this FFN's input, accumulated over a
631    /// calibration run (`CMF_RMS_TRACE`). These are the RMS activation
632    /// traces AWNP needs: raw weight magnitude says every channel matters
633    /// equally, and the question AWNP asks is whether the ACTIVATIONS
634    /// disagree. Off unless the env var is set — an f64 add per channel
635    /// per token is cheap, but not free.
636    pub act_sq: std::cell::RefCell<Vec<f64>>,
637    /// Raw FFN-input rows captured for the layers named by `CMF_ACT_DUMP`
638    /// (`"9,19"`). AWNP is nullspace PROJECTION: after dropping channels the
639    /// survivors are refitted to absorb what was removed, and how much they
640    /// can absorb depends on the activation COVARIANCE, not on per-channel
641    /// RMS. Per-channel numbers can only bound the cost from above.
642    pub act_rows: std::cell::RefCell<Vec<f32>>,
643    /// Task mask over routed experts (DTG-MA over MoE, claim-12 B-field
644    /// applied): `false` experts are excluded from selection, the
645    /// softmax renormalizes over the allowed set. Built by the loader
646    /// from CMF_MOE_MASK=<stats.json> + CMF_MOE_MASK_COVER. None = all.
647    pub mask: Option<Vec<bool>>,
648    /// Gemma-4: per-expert weight scale applied AFTER the top-k renorm
649    /// (`router.per_expert_scale`). None = 1.0 everywhere.
650    pub per_expert_scale: Option<Vec<f32>>,
651    /// Gemma-4: the router reads a SCALE-LESS rms-norm of its input
652    /// (the constant gain router.scale·√hidden is folded into the
653    /// router weights at convert time).
654    pub router_input_norm: bool,
655    /// Cortiq Embryo: resonance routing (P1) — the "logits" are
656    /// bias_e − ‖(x−μ_e) − U_eᵀU_e(x−μ_e)‖², argmax = the expert whose
657    /// descriptor reconstructs the input best. `router` is a placeholder.
658    pub resonance: Option<Resonance>,
659    /// Growth records (`kind = "expert_append"`, spec §9.5.1) mounted
660    /// behind the trunk experts, one entry per grown expert IN THE ORDER
661    /// they sit in `experts` (the tail `experts[experts.len() - grown.len()..]`).
662    /// The trunk keeps `experts.len() - grown.len()` experts. Empty on a
663    /// gated MoE, on a file without records and under `CMF_GROWTH=off`.
664    pub grown: Vec<GrownExpert>,
665}
666
667/// One grown expert of an `expert_append` record as the loader mounted it.
668#[derive(Debug, Clone, PartialEq, Eq)]
669pub struct GrownExpert {
670    /// The record's skill id.
671    pub record: String,
672    /// Position of the record in `header.skills`.
673    pub record_index: usize,
674    pub layer: usize,
675    /// The index the tensor name declares (`experts.{e}` — the chain rule
676    /// of the format); the executed position may be smaller when an
677    /// earlier record of the layer is not mounted.
678    pub expert: usize,
679}
680
681/// Per-expert resonance descriptors of one MoE layer (`mlp.desc.*`).
682pub struct Resonance {
683    /// [E, hidden]
684    pub mu: Vec<f32>,
685    /// [E, k, hidden] orthonormal directions (k may be 0)
686    pub u: Vec<f32>,
687    pub k: usize,
688    /// [E] selection bias (loss-free balancing, trained online)
689    pub bias: Vec<f32>,
690    /// [E] reconstruction-error shell: an expert whose error
691    /// `‖(x−μ)⊥U‖² = d² − proj` exceeds its shell scores `−∞` (it never
692    /// wins). `+inf` = no shell — every trunk expert; a grown expert
693    /// (`expert_append`) carries the finite `desc.shell` its record stores.
694    /// Empty = no shell anywhere (legacy constructors).
695    pub shell: Vec<f32>,
696}
697
698/// `CMF_GROWTH_SHELL` state: 0 = not yet read from the environment, 1 = on,
699/// 2 = off. Process-wide, like the environment it mirrors.
700static GROWTH_SHELL: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
701
702/// Is the growth shell applied (`Resonance::scores` −∞ rule, the resident
703/// graph's packed shell)? `CMF_GROWTH_SHELL=off` disables it for
704/// measurement; [`set_growth_shell`] overrides the environment in-process
705/// (`growth-eval --shell`). Default: on.
706pub fn growth_shell_enabled() -> bool {
707    use std::sync::atomic::Ordering;
708    match GROWTH_SHELL.load(Ordering::Relaxed) {
709        1 => true,
710        2 => false,
711        _ => {
712            let off = std::env::var("CMF_GROWTH_SHELL")
713                .map(|v| v.eq_ignore_ascii_case("off") || v == "0")
714                .unwrap_or(false);
715            GROWTH_SHELL.store(if off { 2 } else { 1 }, Ordering::Relaxed);
716            !off
717        }
718    }
719}
720
721/// Switch the growth shell on/off for this process (`None` = re-read
722/// `CMF_GROWTH_SHELL` on the next query). A pipeline packed into the
723/// resident graph BEFORE the switch keeps the shell it was packed with —
724/// build a new pipeline after switching.
725pub fn set_growth_shell(on: Option<bool>) {
726    GROWTH_SHELL.store(
727        match on {
728            Some(true) => 1,
729            Some(false) => 2,
730            None => 0,
731        },
732        std::sync::atomic::Ordering::Relaxed,
733    );
734}
735
736impl Resonance {
737    /// Does any expert carry a finite shell (a mounted growth record)?
738    pub fn has_shell(&self) -> bool {
739        self.shell.iter().any(|s| s.is_finite())
740    }
741
742    /// The shell the runtime applies right now: the stored one, or all
743    /// `+inf` when the shell is switched off (`CMF_GROWTH_SHELL=off`).
744    pub fn effective_shell(&self, ne: usize) -> Vec<f32> {
745        let mut out = vec![f32::INFINITY; ne];
746        if growth_shell_enabled() {
747            for (o, s) in out.iter_mut().zip(&self.shell) {
748                *o = *s;
749            }
750        }
751        out
752    }
753
754    /// Routing scores for one input row (higher = better). A grown
755    /// expert whose reconstruction error lies outside its shell gets
756    /// `−∞` (unless the shell is switched off); trunk rows are the exact
757    /// bit pattern they were before growth.
758    pub fn scores(&self, x: &[f32], out: &mut [f32]) {
759        let h = x.len();
760        let ne = out.len();
761        let shell_on = growth_shell_enabled() && !self.shell.is_empty();
762        for e in 0..ne {
763            let mu = &self.mu[e * h..(e + 1) * h];
764            let mut d2 = 0.0f32;
765            for j in 0..h {
766                let d = x[j] - mu[j];
767                d2 += d * d;
768            }
769            let mut proj = 0.0f32;
770            for i in 0..self.k {
771                let u = &self.u[(e * self.k + i) * h..(e * self.k + i + 1) * h];
772                let mut p = 0.0f32;
773                for j in 0..h {
774                    p += (x[j] - mu[j]) * u[j];
775                }
776                proj += p * p;
777            }
778            let err = d2 - proj;
779            out[e] = self.bias.get(e).copied().unwrap_or(0.0) - err;
780            if shell_on && err > self.shell.get(e).copied().unwrap_or(f32::INFINITY) {
781                out[e] = f32::NEG_INFINITY;
782            }
783        }
784    }
785}
786
787/// Attention operator of a layer. Extension point: new operators are
788/// new variants here + a forward in their own module.
789pub enum AttnKind {
790    /// GQA softmax attention (+ optional Qwen3.5 qk-norm / output gate).
791    Full {
792        wq: QTensor,
793        wk: QTensor,
794        wv: QTensor,
795        wo: QTensor,
796        q_norm: Option<Vec<f32>>,
797        k_norm: Option<Vec<f32>>,
798        output_gate: bool,
799        /// Laguna: a separate softplus projection applied to the attention
800        /// output before O. The bool means one scalar per head (broadcast
801        /// across head_dim); false means one scalar per element.
802        softplus_gate: Option<(QTensor, bool)>,
803        /// Qwen2-family projection biases (q, k, v).
804        bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
805    },
806    /// Canonical linear core (VMF phase attention).
807    Linear(VmfPhaseWeights),
808    /// Faithful vendor linear operator (Qwen3.5 GatedDeltaNet).
809    LinearGdn(GdnWeights),
810    /// LFM2 gated short-convolution mixer (no KV cache; conv ring state
811    /// lives in the layer's `linear_state`).
812    ShortConv(ShortConvWeights),
813    /// DeepSeek-V2 Multi-head Latent Attention. v1 executes it as
814    /// expand-to-MHA: the latent is projected per token, K/V expand to
815    /// every head and live in the ordinary cache (K head layout
816    /// [rope | nope] so the standard partial rotary covers the shared
817    /// rope key; V rows are zero-padded to the K head_dim and the pad
818    /// is sliced off before O). Latent-resident cache is a later
819    /// optimization, not a semantic change.
820    Mla(Box<MlaWeights>),
821    /// Kimi Delta Attention (Kimi Linear / Kimi-K3): per-channel decayed
822    /// delta rule, separate q/k/v short convs, sigmoid-gated output norm.
823    /// State lives in the layer's `linear_state` (no KV cache).
824    Kda(Box<crate::linear_core::KdaWeights>),
825    /// Natively bounded softmax anchor `swa_sink_v1` (Embryo-O1): ring of
826    /// the last W raw keys with relative RoPE + trained NoPE sinks, one
827    /// softmax. State is the fixed-size ring in `LayerKvCache::bounded`;
828    /// nothing is stored per position (see `crate::bounded`).
829    Bounded(Box<crate::bounded::BoundedWeights>),
830}
831
832/// DeepSeek-V2 MLA projections (see `AttnKind::Mla`).
833pub struct MlaWeights {
834    /// `[nh·(rope+nope), hidden]` (or `[…, q_lora]` when compressed) —
835    /// the converter permutes each head rope-first so rotary_dim =
836    /// qk_rope works unchanged.
837    pub q_proj: QTensor,
838    /// Compressed q (K3/V3 class): x → q_a `[q_lora, hidden]` →
839    /// rms(q_a_norm) → q_proj (= q_b). None = direct q (V2-Lite).
840    pub q_a: Option<QTensor>,
841    pub q_a_norm: Option<Vec<f32>>,
842    /// `kv_a_proj_with_mqa` `[lora + rope, hidden]` (latent first).
843    pub kv_a: QTensor,
844    /// RMS-norm weights over the latent (`kv_a_layernorm`, [lora]).
845    pub kv_a_norm: Vec<f32>,
846    /// `[nh·(nope+v), lora]` — per head [k_nope | v].
847    pub kv_b: QTensor,
848    /// `[hidden, nh·v]`.
849    pub o_proj: QTensor,
850    pub nh: usize,
851    pub qk_rope: usize,
852    pub qk_nope: usize,
853    pub v_dim: usize,
854    pub lora: usize,
855    /// Softmax scale (1/√(rope+nope), YaRN-mscale-corrected at load).
856    pub scale: f32,
857    /// Kimi Linear NoPE: skip the rotary entirely (layout unchanged).
858    pub nope: bool,
859}
860
861/// Multi-token-prediction head (DeepSeek/Qwen style, spec §2.1):
862/// `x = eh_proj·[enorm(embed(next)); hnorm(hidden)]` → one transformer
863/// block over its own KV → shared lm_head. Drafts the token after next;
864/// the main model verifies, so output is exact — MTP only buys speed.
865pub struct MtpModule {
866    pub enorm: Vec<f32>,
867    pub hnorm: Vec<f32>,
868    /// [hidden, 2·hidden]
869    pub eh_proj: QTensor,
870    pub layer: LayerWeights,
871    pub final_norm: Vec<f32>,
872    pub kv: crate::kv_cache::LayerKvCache,
873}
874
875/// A Metal verify graph after its sync: what the commit needs — the
876/// graph (per-layer replay scratch), the GDN layers in encode order (their
877/// CPU states receive the replay), and the attention layers with the CPU
878/// row count they were encoded against (the accepted rows are pulled from
879/// the mirror from there).
880/// One item of the Metal rows-graph plan.
881#[cfg(target_os = "macos")]
882enum MetalRowsItem<'a> {
883    Gdn {
884        run: Vec<crate::gpu_metal::GdnGpuLayer<'a>>,
885        first: usize,
886    },
887    Attn {
888        l: crate::gpu_metal::AttnGpuLayer<'a>,
889        li: usize,
890        q_norm: Option<&'a [f32]>,
891        k_norm: Option<&'a [f32]>,
892        output_gate: bool,
893    },
894}
895
896#[cfg(target_os = "macos")]
897struct MetalVerifyPending {
898    graph: crate::gpu_metal::VerifyGraph,
899    gdn_layers: Vec<usize>,
900    attn_layers: Vec<(usize, usize)>,
901}
902
903/// A round's batched MTP warm-up, submitted but not yet waited
904/// (`mtp_warm_batch_submit` → `mtp_warm_batch_finish`): the trunk commit's
905/// GDN replay is queued between the two.
906#[cfg(target_os = "macos")]
907struct MetalWarmPending {
908    graph: crate::gpu_metal::VerifyGraph,
909    cpu_stored: usize,
910    b: usize,
911}
912
913#[cfg(target_os = "macos")]
914enum MetalRowsRun {
915    /// Capability/preflight refusal before a command buffer was committed.
916    Declined,
917    /// A graph was admitted and then failed; callers must clear the sequence
918    /// rather than replaying it through CPU/serial state.
919    Failed,
920    Completed(MetalVerifyPending),
921}
922
923#[cfg(target_os = "macos")]
924enum MetalPrefillOutcome {
925    Declined,
926    Failed,
927    Completed(Vec<f32>),
928}
929
930#[cfg(target_os = "macos")]
931enum MetalBatchNllOutcome {
932    Declined,
933    Failed(String),
934    Completed(f64, usize),
935}
936
937/// The speculation trial's phases (see the decode loop): four timed
938/// speculative rounds, eight timed plain tokens, then the faster arm
939/// until a re-check.
940#[derive(Clone, Copy)]
941enum SpecTrial {
942    Spec {
943        t0: std::time::Instant,
944        gen0: usize,
945        rounds: usize,
946    },
947    Plain {
948        t0: std::time::Instant,
949        gen0: usize,
950    },
951    Decided {
952        spec: bool,
953        recheck_at: usize,
954    },
955}
956
957/// `CMF_GRAPH_SPEC_TIME`: 0 = off, 1 = one line per speculative round
958/// plus the host stamps of any OUTLIER round (wall > 1.4× the running
959/// median), 2 = the host stamps of every round.
960pub(crate) fn spec_time_level() -> u8 {
961    static L: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
962    *L.get_or_init(|| match std::env::var("CMF_GRAPH_SPEC_TIME") {
963        Ok(v) => v.trim().parse::<u8>().map(|n| n.max(1)).unwrap_or(1),
964        Err(_) => 0,
965    })
966}
967
968/// The round's host stamps: `spec_stamp(name)` records the time since
969/// the previous stamp (the section that just ended) — from anywhere on
970/// the round's call chain (the Metal verify, the draft step, the commit),
971/// no plumbing. Off (a single atomic load) unless `CMF_GRAPH_SPEC_TIME`
972/// is set; one decode thread at a time is assumed (diagnostics).
973struct SpecStampLog {
974    t_last: std::time::Instant,
975    items: Vec<(&'static str, f32)>,
976}
977
978static SPEC_STAMPS: std::sync::Mutex<Option<SpecStampLog>> = std::sync::Mutex::new(None);
979
980pub(crate) fn spec_stamp(name: &'static str) {
981    if spec_time_level() == 0 {
982        return;
983    }
984    if let Ok(mut g) = SPEC_STAMPS.lock() {
985        if let Some(log) = g.as_mut() {
986            let now = std::time::Instant::now();
987            log.items
988                .push((name, (now - log.t_last).as_secs_f32() * 1e3));
989            log.t_last = now;
990        }
991    }
992}
993
994fn spec_stamps_begin() {
995    if spec_time_level() == 0 {
996        return;
997    }
998    if let Ok(mut g) = SPEC_STAMPS.lock() {
999        *g = Some(SpecStampLog {
1000            t_last: std::time::Instant::now(),
1001            items: Vec::with_capacity(64),
1002        });
1003    }
1004}
1005
1006fn spec_stamps_take() -> Vec<(&'static str, f32)> {
1007    SPEC_STAMPS
1008        .lock()
1009        .ok()
1010        .and_then(|mut g| g.take())
1011        .map(|l| l.items)
1012        .unwrap_or_default()
1013}
1014
1015/// One line: every stamp name in first-seen order with its total over the
1016/// round and, when it fired more than once (the draft steps), the count.
1017fn spec_stamps_format(items: &[(&'static str, f32)]) -> String {
1018    let mut agg: Vec<(&'static str, f32, u32)> = Vec::with_capacity(items.len());
1019    for &(n, ms) in items {
1020        match agg.iter_mut().find(|e| e.0 == n) {
1021            Some(e) => {
1022                e.1 += ms;
1023                e.2 += 1;
1024            }
1025            None => agg.push((n, ms, 1)),
1026        }
1027    }
1028    let mut s = String::with_capacity(agg.len() * 16);
1029    for (n, ms, k) in agg {
1030        if k > 1 {
1031            s.push_str(&format!("{n} {ms:.1}/{k} "));
1032        } else {
1033            s.push_str(&format!("{n} {ms:.1} "));
1034        }
1035    }
1036    s
1037}
1038
1039/// The speculation monitor: exponential averages of a round's wall time
1040/// and of the tokens it produced, and the plain token's wall time — the
1041/// three numbers the keep/stop rule needs. A round pays when
1042/// `tokens_per_round · plain_ms > round_ms · 1.03`. The one-shot trial
1043/// (four rounds against eight tokens) mis-called prose: the first rounds
1044/// after a prompt are formulaic and accept well, the body does not (an
1045/// essay measured 39 against a plain 44.8 with the trial saying
1046/// "speculate"), so the rule now runs on EVERY round and stops after four
1047/// consecutive losing rounds; a stopped speculation is retried 128 tokens
1048/// later.
1049///
1050/// Native Metal (`metal: true`) does not pay the eight plain tokens up
1051/// front: on the 27B a plain token is ~150 ms, so the trial alone cost
1052/// ~1.2 s of every answer. There the plain phase is (a) skipped while the
1053/// rounds land at least `SPEC_PROXY_TOKENS` tokens each — a k=7 round on
1054/// Metal costs ~1.9 plain tokens (286 against 148 ms measured on the M4),
1055/// so 3.5 tokens/round cannot lose on any Metal round/plain ratio seen —
1056/// and (b) otherwise bounded to the fewest tokens that time it: two, or
1057/// as many as fit in `SPEC_PLAIN_MIN_MS` (a 150-ms token measures itself;
1058/// a 10-ms one needs the eight). The keep/stop rule itself is unchanged:
1059/// the moment a plain rate exists, it decides.
1060#[derive(Default, Clone, Copy)]
1061struct SpecMon {
1062    round_ms: f64,
1063    tokens: f64,
1064    plain_ms: f64,
1065    n: u32,
1066    fails: u32,
1067    metal: bool,
1068}
1069
1070/// Tokens per round at or above which a Metal round pays without a plain
1071/// measurement (see `SpecMon`).
1072const SPEC_PROXY_TOKENS: f64 = 3.5;
1073/// The Metal plain phase: at least two tokens, and more until this much
1074/// wall time has been timed (up to the eight the other backends time).
1075const SPEC_PLAIN_MIN_MS: f64 = 200.0;
1076
1077impl SpecMon {
1078    fn round(&mut self, dt_ms: f64, produced: usize) {
1079        self.n += 1;
1080        if self.n == 1 {
1081            return; // round 1 pays the batch scratch and the draft mirror
1082        }
1083        let a = if self.n == 2 { 1.0 } else { 0.3 };
1084        self.round_ms += a * (dt_ms - self.round_ms);
1085        self.tokens += a * (produced as f64 - self.tokens);
1086    }
1087    fn pays(&self) -> bool {
1088        if self.plain_ms > 0.0 {
1089            self.tokens * self.plain_ms > self.round_ms * 1.03
1090        } else {
1091            self.metal && self.tokens >= SPEC_PROXY_TOKENS
1092        }
1093    }
1094    /// Has the plain phase timed enough tokens to decide?
1095    fn plain_done(&self, t0: std::time::Instant, gen0: usize, generated: usize) -> bool {
1096        let n = generated.saturating_sub(gen0);
1097        if n >= 8 {
1098            return true;
1099        }
1100        self.metal && n >= 2 && t0.elapsed().as_secs_f64() * 1e3 >= SPEC_PLAIN_MIN_MS
1101    }
1102}
1103
1104/// Ids of the consumed prefix the bounded reuse key remembers literally
1105/// (the rest is covered by the rolling hash).
1106pub const KV_PREFIX_TAIL: usize = 128;
1107
1108/// Bounded prefix-reuse key: how many ids the cache holds, a rolling
1109/// hash of ALL of them and the last [`KV_PREFIX_TAIL`] ids literally.
1110/// Answers "does this prompt strictly extend what is cached" exactly
1111/// (hash over the whole consumed prefix + literal tail) without keeping
1112/// the dialogue — the record is a fixed size whatever the session length.
1113#[derive(Debug, Clone, Default)]
1114pub struct KvPrefix {
1115    len: usize,
1116    hash: u64,
1117    tail: Vec<u32>,
1118    /// Owner of the state this key describes: the resident device graph
1119    /// (true) or the host. A turn continues the prefix only on its owner
1120    /// (R4: a host continuation of a device sequence reads an empty host
1121    /// state; a device continuation of a host sequence has no image).
1122    device: bool,
1123}
1124
1125impl KvPrefix {
1126    #[inline]
1127    fn fold(mut h: u64, ids: &[u32]) -> u64 {
1128        for &id in ids {
1129            h ^= id as u64;
1130            h = h.wrapping_mul(0x100000001b3);
1131            h ^= h >> 29;
1132        }
1133        h
1134    }
1135
1136    pub fn clear(&mut self) {
1137        self.len = 0;
1138        self.hash = 0xcbf29ce484222325;
1139        self.tail.clear();
1140        self.device = false;
1141    }
1142
1143    /// Was the prefix built on the resident device graph?
1144    pub fn on_device(&self) -> bool {
1145        self.device
1146    }
1147
1148    /// Tag the owner of the recorded prefix.
1149    pub fn set_on_device(&mut self, device: bool) {
1150        self.device = device;
1151    }
1152
1153    /// Ids the cache holds (the forwarded prefix).
1154    pub fn len(&self) -> usize {
1155        self.len
1156    }
1157
1158    pub fn is_empty(&self) -> bool {
1159        self.len == 0
1160    }
1161
1162    /// Literal tail currently kept (≤ `KV_PREFIX_TAIL`).
1163    pub fn tail_len(&self) -> usize {
1164        self.tail.len()
1165    }
1166
1167    /// Replace the key with `ids` (a fresh sequence).
1168    pub fn set(&mut self, ids: &[u32]) {
1169        self.clear();
1170        self.extend(ids);
1171    }
1172
1173    /// Append `more` to the consumed prefix (an extension-only turn).
1174    pub fn extend(&mut self, more: &[u32]) {
1175        if self.len == 0 && self.hash == 0 {
1176            self.hash = 0xcbf29ce484222325;
1177        }
1178        self.hash = Self::fold(self.hash, more);
1179        self.len += more.len();
1180        if more.len() >= KV_PREFIX_TAIL {
1181            self.tail.clear();
1182            self.tail.extend_from_slice(&more[more.len() - KV_PREFIX_TAIL..]);
1183        } else {
1184            let drop = (self.tail.len() + more.len()).saturating_sub(KV_PREFIX_TAIL);
1185            self.tail.drain(..drop);
1186            self.tail.extend_from_slice(more);
1187        }
1188    }
1189
1190    /// Cached positions when `ids` strictly extends the consumed prefix,
1191    /// 0 otherwise. The tail is compared literally first (cheap), then
1192    /// the hash over the whole prefix must agree.
1193    pub fn extension(&self, ids: &[u32]) -> usize {
1194        if self.len == 0 || ids.len() <= self.len {
1195            return 0;
1196        }
1197        let t = self.tail.len();
1198        if ids[self.len - t..self.len] != self.tail[..] {
1199            return 0;
1200        }
1201        if Self::fold(0xcbf29ce484222325, &ids[..self.len]) != self.hash {
1202            return 0;
1203        }
1204        self.len
1205    }
1206}
1207
1208/// Result of a generation call.
1209pub struct GenerateResult {
1210    pub text: String,
1211    pub token_ids: Vec<u32>,
1212    pub prompt_tokens: usize,
1213    pub tokens_generated: usize,
1214    pub finish_reason: String,
1215    /// Speculative-decode stats (0/0 when MTP is absent or inactive).
1216    pub mtp_drafted: usize,
1217    pub mtp_accepted: usize,
1218    /// Per-generated-token confidence = softmax probability of the token
1219    /// that was actually emitted (softmax probability on the chosen state). High =
1220    /// the model was sure; low = it was guessing. Same length as the
1221    /// generated slice of `token_ids`.
1222    pub token_confidence: Vec<f32>,
1223    /// Structured per-token telemetry (B4 channel). Empty unless
1224    /// `set_trace(true)`; otherwise same length as the generated slice.
1225    pub traces: Vec<TokenTrace>,
1226}
1227
1228/// One row of the structured telemetry trace (B4): the model's internal
1229/// routing state at the moment a token was emitted. Every field is a
1230/// quantity the runtime already computes — nothing is inferred or
1231/// estimated (anti-principle: only measured bytes).
1232#[derive(Clone, Debug)]
1233pub struct TokenTrace {
1234    /// 0-based index within the generated slice.
1235    pub t: usize,
1236    /// The emitted token id.
1237    pub token_id: u32,
1238    /// Softmax probability on the emitted token — how sure the model was.
1239    pub confidence: f32,
1240    /// Skill in force while this token was generated (None = backbone).
1241    pub active_skill: Option<String>,
1242    /// Recon error E = ‖r−BBᵀr‖²/‖φ‖² at the last routing eval — coherence
1243    /// with the active skill's subspace (low = coherent). None = no router
1244    /// or not yet evaluated.
1245    pub recon: Option<f32>,
1246    /// The router changed the active skill right after this token (a
1247    /// domain boundary crossed under the hysteresis barrier).
1248    pub switched: bool,
1249}
1250
1251/// Calibrated softmax probability of `id` under `logits` (the confidence on
1252/// the emitted token) — the confidence signal, cheap from logits already
1253/// computed for sampling. `temp` is the calibration temperature (B1):
1254/// softmax(logits / temp); 1.0 = raw.
1255#[cfg_attr(not(test), allow(dead_code))]
1256fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
1257    let t = if temp > 1e-3 { temp } else { 1.0 };
1258    let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1259    let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
1260    if sum > 0.0 {
1261        (((logits[id as usize] - max) / t).exp()) / sum
1262    } else {
1263        0.0
1264    }
1265}
1266
1267/// prefill-GEMM enabled? (CMF_PREFILL=seq — emergency fallback to the
1268/// sequential path.)
1269fn prefill_batched() -> bool {
1270    std::env::var("CMF_PREFILL")
1271        .map(|v| v != "seq")
1272        .unwrap_or(true)
1273}
1274
1275/// Decide the graph NLL route without conflating graph quality with the
1276/// optional native-Metal fused head. A hidden-state graph remains a valid
1277/// quality route on Vulkan/Wgpu; only native Metal requires graph logits.
1278#[inline]
1279fn nll_graph_policy(
1280    unmasked: bool,
1281    prefer_graph: bool,
1282    native_metal: bool,
1283) -> (bool, bool) {
1284    let graph_quality = unmasked && prefer_graph;
1285    let fused_head_quality = graph_quality && native_metal;
1286    (graph_quality, fused_head_quality)
1287}
1288
1289/// Input to the layer-major batched span walk: token ids (embeds itself,
1290/// full-stack and coordinator prefill) or ready boundary hiddens (the
1291/// network worker's side of a split).
1292#[derive(Clone, Copy)]
1293enum PrefillIn<'a> {
1294    Ids(&'a [u32]),
1295    Hidden(&'a [f32]),
1296}
1297
1298/// The batched prefill walks `weights.layers`. Architectures that load
1299/// their own stack (gemma-3n's AltUp replicas, DeepSeek-V4's hyper-
1300/// connections) leave that empty and must go position by position — asking
1301/// otherwise indexes an empty vector, which is a panic rather than a
1302/// fallback. Every call site goes through here so the next such
1303/// architecture is one line, not four.
1304impl Pipeline {
1305    fn can_prefill_batched(&self) -> bool {
1306        #[cfg(test)]
1307        let force_serial = self.nll_test_force_serial;
1308        #[cfg(not(test))]
1309        let force_serial = false;
1310        prefill_batched() && !force_serial && !self.weights.layers.is_empty()
1311    }
1312
1313    /// The backend's automatic capacity split for a mapped transformer.
1314    /// Kept as a method so prefill and decode use the exact same boundary.
1315    fn automatic_gpu_prefix(&self) -> Option<usize> {
1316        let (model, _, _, _) = self.weights.embed_tokens.graph_weight()?;
1317        crate::gpu::automatic_layer_prefix(&model, self.num_layers, self.physical_layers)
1318    }
1319
1320    /// Positions per batched pass of the layer-stack prefill for THIS
1321    /// model on THIS backend (see [`prefill_chunk_rule`]). Pub: the network
1322    /// split must chunk exactly like the local path to reproduce it.
1323    pub fn prefill_chunk(&self) -> usize {
1324        let env = env_prefill_chunk();
1325        if env.is_some() || ChunkHost::here() != ChunkHost::Other {
1326            return prefill_chunk_rule(env, ChunkHost::here(), false);
1327        }
1328        prefill_chunk_rule(None, ChunkHost::Other, self.chunk_stack_facts().dense_on_discrete())
1329    }
1330
1331    fn chunk_stack_facts(&self) -> ChunkStackFacts {
1332        let plain_dense = !self.weights.layers.is_empty()
1333            && self.g3n.is_none()
1334            && self.dsv4.is_none()
1335            && self.dsv41.is_none()
1336            && self.qwen4_exp.is_none()
1337            && self.weights.layers.iter().all(|lw| {
1338                matches!(lw.attn, AttnKind::Full { .. }) && matches!(lw.ffn, FfnKind::Dense(_))
1339            });
1340        let gpu_on = crate::gpu::enabled();
1341        ChunkStackFacts {
1342            plain_dense,
1343            discrete: gpu_on && crate::gpu::discrete(),
1344            gpu_on,
1345            // Only asked when the rest already qualifies: it opens the
1346            // backend's capacity plan.
1347            capacity_split: std::env::var_os("CMF_GPU_LAYERS").is_some()
1348                || (plain_dense && gpu_on && self.automatic_gpu_prefix().is_some()),
1349            multi_gpu: self.gpu_plan.is_some(),
1350            o1: self.o1_active(),
1351        }
1352    }
1353}
1354
1355/// Prefill chunk (positions per batched pass), model-agnostic form. On
1356/// macOS the AMX GEMM path wants tall panels — M=48 starves the matrix
1357/// units (ggml uses ubatch 512); elsewhere the historical 48 stays.
1358/// CMF_PREFILL_CHUNK overrides. The architectures with their own stacks
1359/// (DeepSeek-V4/V4.1) chunk with this; the layer-stack prefill asks
1360/// [`Pipeline::prefill_chunk`], which also knows the model and the card.
1361/// A different chunk is a different (equally valid) generation: panel
1362/// width reorders float accumulation.
1363pub fn prefill_chunk() -> usize {
1364    prefill_chunk_rule(env_prefill_chunk(), ChunkHost::here(), false)
1365}
1366
1367fn env_prefill_chunk() -> Option<usize> {
1368    std::env::var("CMF_PREFILL_CHUNK")
1369        .ok()
1370        .and_then(|v| v.parse::<usize>().ok())
1371}
1372
1373/// The host classes the chunk width distinguishes.
1374#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1375enum ChunkHost {
1376    Macos,
1377    /// Linux/Android aarch64 (phones, SBCs).
1378    Aarch64,
1379    /// Everything else: x86-64 Linux/Windows, CPU or Vulkan/DX12.
1380    Other,
1381}
1382
1383impl ChunkHost {
1384    fn here() -> Self {
1385        if cfg!(target_os = "macos") {
1386            ChunkHost::Macos
1387        } else if cfg!(target_arch = "aarch64") {
1388            ChunkHost::Aarch64
1389        } else {
1390            ChunkHost::Other
1391        }
1392    }
1393}
1394
1395/// Chunk for a plain dense stack whose every layer lives on a discrete
1396/// card. On x86 the layer-stack prefill is host-driven: each GEMM and the
1397/// chunk attention (which re-uploads the whole KV prefix per layer) is a
1398/// separate submit + readback, so 48 positions a pass left the card idle
1399/// between them. Measured in-process on an RTX 3090 (Vulkan), 2048-token
1400/// prompt — see CHANGELOG 0.7.6 for the table.
1401const DISCRETE_DENSE_PREFILL_CHUNK: usize = 512;
1402
1403/// The chunk-width rule. `dense_on_discrete` is true only for a plain
1404/// dense transformer (full attention, dense FFN, no special stack) that
1405/// is entirely resident on one discrete card — the one case measured
1406/// here. GDN hybrids, MoE, DeepSeek stacks, capacity-split and CPU-only
1407/// runs keep the width they were tuned with.
1408fn prefill_chunk_rule(env: Option<usize>, host: ChunkHost, dense_on_discrete: bool) -> usize {
1409    if let Some(n) = env {
1410        return n.max(1);
1411    }
1412    match host {
1413        ChunkHost::Macos => 512,
1414        // Mobile: big enough to feed the batched attend (gate b ≥ 32)
1415        // and the blocked SDOT GEMM without the memory of 512.
1416        ChunkHost::Aarch64 => 256,
1417        ChunkHost::Other if dense_on_discrete => DISCRETE_DENSE_PREFILL_CHUNK,
1418        ChunkHost::Other => 48,
1419    }
1420}
1421
1422/// What the chunk rule needs to know about a loaded stack.
1423#[derive(Clone, Copy, Debug, Default)]
1424struct ChunkStackFacts {
1425    /// Every layer is `AttnKind::Full` + `FfnKind::Dense`, and no
1426    /// architecture-owned stack (g3n, DeepSeek-V4/V4.1, qwen4-exp) is set.
1427    plain_dense: bool,
1428    /// The active GPU backend is a discrete card.
1429    discrete: bool,
1430    /// The backend is up and not paused.
1431    gpu_on: bool,
1432    /// A capacity-derived device prefix: some layers run on the host.
1433    capacity_split: bool,
1434    /// An in-process multi-GPU plan is set.
1435    multi_gpu: bool,
1436    /// O(1) layers (their Q trace is recorded by the prefill).
1437    o1: bool,
1438}
1439
1440impl ChunkStackFacts {
1441    fn dense_on_discrete(self) -> bool {
1442        self.plain_dense
1443            && self.discrete
1444            && self.gpu_on
1445            && !self.capacity_split
1446            && !self.multi_gpu
1447            && !self.o1
1448    }
1449}
1450
1451/// Number of prompt rows that have a real teacher-forced next-token pair in a
1452/// prefill span.  The final prompt row has no successor token, so it must not
1453/// be handed to the MTP warm-up.  Keeping this arithmetic in one helper makes
1454/// the full-chunk and tail-chunk boundaries explicit for both the graph and
1455/// CPU implementations.
1456#[inline]
1457fn mtp_prefill_pair_count(start: usize, end: usize, input_len: usize) -> usize {
1458    if end <= start || start >= input_len {
1459        return 0;
1460    }
1461    let rows = (end.min(input_len) - start).min(input_len - start);
1462    if end < input_len {
1463        rows
1464    } else {
1465        rows.saturating_sub(1)
1466    }
1467}
1468
1469/// Callback for streaming tokens. Return `false` to cancel.
1470pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
1471
1472/// One layer's cache ownership at a cross-turn KV reuse boundary.
1473#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1474pub(crate) struct ReuseLayer {
1475    /// Exact-attention layer (rows in `LayerKvCache`); otherwise a
1476    /// recurrent / latent mixer whose state cannot be rewound.
1477    pub full: bool,
1478    /// Rows the host owner cache holds.
1479    pub host_rows: usize,
1480    /// Rows the wgpu token graph's device mirror holds (None: no mirror).
1481    pub device_rows: Option<usize>,
1482    /// A recurrent state lives on the device (advanced past the host copy).
1483    pub device_state: bool,
1484}
1485
1486/// What a reused turn must do before its tail prefill runs on the HOST.
1487#[derive(Debug, Clone, PartialEq, Eq)]
1488pub(crate) enum ReusePlan {
1489    /// Host caches already hold exactly the reused prefix.
1490    Ready,
1491    /// Copy device mirror rows `[from..to)` into the host cache of each
1492    /// listed layer (the rows decode wrote on the device only).
1493    Pull(Vec<(usize, usize, usize)>),
1494    /// The prefix cannot be continued on the host exactly: start fresh.
1495    Fresh,
1496}
1497
1498/// The wgpu whole-token graph decodes into a DEVICE K/V mirror and never
1499/// writes those rows back to the host cache, while the chunked prefill of a
1500/// pure-attention model reads (and appends to) the host cache. A reused turn
1501/// therefore found its host cache ending at the previous PROMPT, not at the
1502/// previous answer: the tail prefill attended without the model's own
1503/// answer and appended its rows at the wrong index (MiniCPM5 on Vulkan
1504/// repeated its tool call instead of reading the tool result). Every layer
1505/// must hold exactly `reuse_from` host rows before the host continues; rows
1506/// that exist only on the device are pulled back, anything else is fresh.
1507pub(crate) fn kv_reuse_plan(reuse_from: usize, layers: &[ReuseLayer]) -> ReusePlan {
1508    let mut pulls = Vec::new();
1509    for (li, l) in layers.iter().enumerate() {
1510        if !l.full {
1511            if l.device_state {
1512                return ReusePlan::Fresh;
1513            }
1514            continue;
1515        }
1516        if l.host_rows == reuse_from {
1517            continue;
1518        }
1519        if l.host_rows < reuse_from && l.device_rows.is_some_and(|d| d >= reuse_from) {
1520            pulls.push((li, l.host_rows, reuse_from));
1521            continue;
1522        }
1523        return ReusePlan::Fresh;
1524    }
1525    if pulls.is_empty() {
1526        ReusePlan::Ready
1527    } else {
1528        ReusePlan::Pull(pulls)
1529    }
1530}
1531
1532impl Pipeline {
1533    /// Clear all per-sequence state, including backend device mirrors.
1534    ///
1535    /// The host KV/history buffers are only half of the request lifecycle on
1536    /// wgpu: GDN/O(1) state and cached graph bind groups are keyed by the
1537    /// pipeline id and otherwise survive a pooled request.  Keep every fresh
1538    /// sequence entry point on this one reset path so a new request cannot
1539    /// inherit the prior request's device state.
1540    fn clear_sequence_state(&mut self) {
1541        // a replay still writing the GDN owners must land before they are
1542        // cleared or reallocated (the device holds raw pointers to them)
1543        #[cfg(target_os = "macos")]
1544        let _ = crate::gpu_metal::wait_replay();
1545        self.kv_cache.clear();
1546        // Both reuse keys (the legacy `kv_history` and the bounded
1547        // `kv_prefix`) describe the state being dropped here.
1548        self.clear_history();
1549        self.graph_logits = None;
1550        if let Some(b) = &mut self.dsv41 {
1551            b.3.clear();
1552        }
1553        crate::gpu::graph_kv_reset(self.graph_kv_id);
1554        // MTP is detached from `self` for the duration of generation, so its
1555        // device mirror is not covered by the trunk reset above.  Reset the
1556        // derived id as well: a failed/aborted warm-up must never leave a
1557        // mirror that a later request can mistake for a current MTP cache.
1558        crate::gpu::graph_kv_reset(self.mtp_kv_id());
1559    }
1560
1561    /// Make the host caches own exactly the reused prefix `[0..reuse_from)`
1562    /// before a reused turn's tail prefill runs on the host (see
1563    /// [`kv_reuse_plan`]). Returns false when the prefix cannot be continued
1564    /// exactly — the caller then starts a fresh sequence. A model whose
1565    /// prefill runs through the token graph keeps its device state as the
1566    /// authority and is left untouched.
1567    fn prepare_kv_reuse(&mut self, reuse_from: usize) -> bool {
1568        if self.graph_prefill_preferred() {
1569            return true;
1570        }
1571        let kv_id = self.graph_kv_id;
1572        let layers: Vec<ReuseLayer> = (0..self.num_layers)
1573            .map(|li| {
1574                let full = matches!(
1575                    self.weights.layers[self.phys_layer(li)].attn,
1576                    AttnKind::Full { .. }
1577                );
1578                ReuseLayer {
1579                    full,
1580                    host_rows: self.kv_cache.layers[li].seq_len,
1581                    device_rows: crate::gpu::graph_kv_stored(kv_id, li),
1582                    device_state: crate::gpu::graph_state_resident(kv_id, li),
1583                }
1584            })
1585            .collect();
1586        // No wgpu device state at all (CPU, Metal — whose graph appends every
1587        // decoded row to the owner cache itself): the host is the owner and
1588        // the extension check already proved the prefix.
1589        if layers
1590            .iter()
1591            .all(|l| l.device_rows.is_none() && !l.device_state)
1592        {
1593            return true;
1594        }
1595        let plan = kv_reuse_plan(reuse_from, &layers);
1596        let (what, rows, n) = match &plan {
1597            ReusePlan::Ready => ("host ready", 0, 0),
1598            ReusePlan::Fresh => ("fresh", 0, 0),
1599            ReusePlan::Pull(p) => (
1600                "pull",
1601                p.iter().map(|&(_, a, b)| b - a).max().unwrap_or(0),
1602                p.len(),
1603            ),
1604        };
1605        let t0 = std::time::Instant::now();
1606        let ok = self.apply_kv_reuse_plan(reuse_from, plan, &layers);
1607        if std::env::var("CMF_PREFILL_PROF").is_ok() {
1608            eprintln!(
1609                "kv-reuse: {what}{}: {rows} device row(s) × {n} layer(s) to the host in {:.2} ms",
1610                if ok { "" } else { " (failed → fresh)" },
1611                t0.elapsed().as_secs_f64() * 1e3
1612            );
1613        }
1614        ok
1615    }
1616
1617    fn apply_kv_reuse_plan(
1618        &mut self,
1619        reuse_from: usize,
1620        plan: ReusePlan,
1621        layers: &[ReuseLayer],
1622    ) -> bool {
1623        let kv_id = self.graph_kv_id;
1624        match plan {
1625            ReusePlan::Fresh => return false,
1626            ReusePlan::Ready => {}
1627            ReusePlan::Pull(pulls) => {
1628                // Mirrors of one uniform geometry: one batched read serves
1629                // every layer. Per-layer geometry (MiMo-V2: 4/8 KV heads,
1630                // narrow V, sliding rings) is read layer by layer in the
1631                // host layout instead.
1632                let (nkv, hd) = {
1633                    let c = &self.kv_cache.layers[pulls[0].0];
1634                    (c.num_kv_heads, c.head_dim)
1635                };
1636                let uniform = pulls.iter().all(|&(li, _, _)| {
1637                    let c = &self.kv_cache.layers[li];
1638                    (c.num_kv_heads, c.head_dim) == (nkv, hd)
1639                });
1640                let batched = if uniform {
1641                    crate::gpu::graph_kv_read_rows(kv_id, &pulls, nkv, hd)
1642                } else {
1643                    None
1644                };
1645                let rows: Vec<(Vec<f32>, Vec<f32>)> = match batched {
1646                    Some(rows) => rows,
1647                    None => {
1648                        let mut rows = Vec::with_capacity(pulls.len());
1649                        for &(li, from, to) in &pulls {
1650                            let (lnkv, lhd) = {
1651                                let c = &self.kv_cache.layers[li];
1652                                (c.num_kv_heads, c.head_dim)
1653                            };
1654                            let Some((k, v, first_valid)) =
1655                                crate::gpu::graph_kv_pull_host(kv_id, li, from, to, lnkv, lhd)
1656                            else {
1657                                return false;
1658                            };
1659                            // The host continues at `to`: a sliding layer
1660                            // reads back only its last window, a full one
1661                            // every row it lacks.
1662                            let need_from = match self.layer_window(li) {
1663                                Some(w) => from.max((to + 1).saturating_sub(w)),
1664                                None => from,
1665                            };
1666                            if first_valid > need_from {
1667                                return false;
1668                            }
1669                            rows.push((k, v));
1670                        }
1671                        rows
1672                    }
1673                };
1674                for ((li, from, to), (k, v)) in pulls.into_iter().zip(rows) {
1675                    let cache = &mut self.kv_cache.layers[li];
1676                    let row = cache.num_kv_heads * cache.head_dim;
1677                    for p in 0..to - from {
1678                        cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
1679                    }
1680                    if cache.seq_len != to {
1681                        return false;
1682                    }
1683                }
1684            }
1685        }
1686        // A mirror past the prefix (a greedy burst that ran beyond the stop)
1687        // holds rows of the OLD continuation: rewind it so the next graph
1688        // token re-syncs those positions from the host.
1689        for (li, l) in layers.iter().enumerate() {
1690            if l.full
1691                && l.device_rows.is_some_and(|d| d > reuse_from)
1692                && !crate::gpu::graph_kv_set_stored(kv_id, li, reuse_from)
1693            {
1694                return false;
1695            }
1696        }
1697        true
1698    }
1699
1700    /// Finish a generation lifecycle after the MTP/router owners were
1701    /// detached.  Every terminal path must put those owners back before the
1702    /// pooled pipeline can serve another request.  Graph side channels and
1703    /// device mirrors are cleared on errors and cancellations; a successful
1704    /// generation keeps its decode-ready host cache for KV reuse.
1705    fn finish_generation(
1706        &mut self,
1707        mtp: &mut Option<MtpModule>,
1708        router: &mut Option<crate::swarm::DynRouter>,
1709        clear_sequence: bool,
1710    ) {
1711        // A dynamic route may have switched the overlay before the terminal
1712        // path. Restore the backbone while the detached router is still
1713        // available, because set_active_skill also owns the overlay reset.
1714        if router.is_some() {
1715            let _ = self.set_active_skill(None);
1716        }
1717        // The last speculative round's replay may still be in flight on
1718        // the second queue: whoever reads the host cache after generate()
1719        // returns (session export, the network split's KV wire, a KV
1720        // reuse) must see the final states.
1721        // A replay that failed leaves the GDN owners half-written: fail
1722        // closed and drop the sequence instead of handing the cache on.
1723        #[cfg(target_os = "macos")]
1724        let clear_sequence = clear_sequence || !crate::gpu_metal::wait_replay();
1725        if clear_sequence {
1726            self.clear_sequence_state();
1727            if let Some(m) = mtp.as_mut() {
1728                // The MTP owner is detached while generation runs, so the
1729                // trunk reset above cannot clear its host cache.  Drop its
1730                // partial rows before reattaching it to the pooled pipeline;
1731                // the next request must start from the same empty anchor on
1732                // CPU and on the device mirror.
1733                m.kv.clear();
1734            }
1735            if let Some(m) = self.mtp.as_mut() {
1736                // A non-speculative request leaves the configured MTP owner
1737                // attached.  Clear that dormant cache too when a shared
1738                // generation failure/cancellation resets the sequence.
1739                m.kv.clear();
1740            }
1741        }
1742        self.graph_want_logits = false;
1743        self.graph_head_required = false;
1744        self.graph_logits = None;
1745        self.graph_failed
1746            .store(false, std::sync::atomic::Ordering::Relaxed);
1747        self.cancel
1748            .store(false, std::sync::atomic::Ordering::Relaxed);
1749        self.dyn_router = router.take().or(self.dyn_router.take());
1750        self.mtp = mtp.take().or(self.mtp.take());
1751        self.mtp_graph_mode = None;
1752        self.spec_forced = None;
1753    }
1754
1755    /// Consume a graph failure reported by a forward that returns only a
1756    /// hidden vector.  `forward_ids` is a public Result API, so it must not
1757    /// turn the graph's zero hidden sentinel into a valid lm_head result.
1758    fn check_forward_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1759        if self
1760            .graph_failed
1761            .swap(false, std::sync::atomic::Ordering::Relaxed)
1762        {
1763            self.cancel
1764                .store(false, std::sync::atomic::Ordering::Relaxed);
1765            self.clear_sequence_state();
1766            self.graph_logits = None;
1767            self.graph_want_logits = false;
1768            self.graph_head_required = false;
1769            return Err(format!("GPU graph failed during {phase} at position {pos}"));
1770        }
1771        Ok(())
1772    }
1773
1774    #[cfg(target_os = "macos")]
1775    fn fail_metal_graph(&mut self, reason: &str) {
1776        crate::pipeline::METAL_GRAPH_ERRORS
1777            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1778        self.clear_sequence_state();
1779        self.graph_logits = None;
1780        self.graph_failed
1781            .store(true, std::sync::atomic::Ordering::Relaxed);
1782        self.cancel
1783            .store(true, std::sync::atomic::Ordering::Relaxed);
1784        tracing::error!("native Metal TokenGraph failed closed: {reason}");
1785    }
1786
1787    /// Start an NLL/PPL request with all graph side channels in a known
1788    /// state.  A graph failure also raises the cooperative cancel bit; it is
1789    /// consumed here and that graph-induced bit is cleared so an independent
1790    /// request can be reused.  A caller-owned cancellation remains intact.
1791    fn nll_begin(&mut self) -> Result<(), String> {
1792        if self
1793            .graph_failed
1794            .swap(false, std::sync::atomic::Ordering::Relaxed)
1795        {
1796            self.cancel
1797                .store(false, std::sync::atomic::Ordering::Relaxed);
1798            self.clear_sequence_state();
1799            self.graph_logits = None;
1800            self.graph_want_logits = false;
1801            self.graph_head_required = false;
1802            return Err("GPU graph failed before NLL scoring".to_string());
1803        }
1804        self.clear_sequence_state();
1805        self.graph_logits = None;
1806        self.graph_want_logits = false;
1807        self.graph_head_required = false;
1808        Ok(())
1809    }
1810
1811    /// End an NLL/PPL request, including the side channels that are not part
1812    /// of the host KV cache.  This is intentionally explicit instead of
1813    /// relying on a tuple/sentinel return: callers must see every failure.
1814    fn nll_end(&mut self) {
1815        self.clear_sequence_state();
1816        self.graph_logits = None;
1817        self.graph_want_logits = false;
1818        self.graph_head_required = false;
1819        self.graph_failed
1820            .store(false, std::sync::atomic::Ordering::Relaxed);
1821    }
1822
1823    /// Check the graph failure channel at a scoring boundary and leave the
1824    /// pipeline reusable when the device path failed.
1825    fn nll_check_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1826        #[cfg(test)]
1827        if self.nll_test_fail_at == Some(pos) {
1828            self.nll_test_fail_at = None;
1829            self.graph_failed
1830                .store(true, std::sync::atomic::Ordering::Relaxed);
1831            self.cancel
1832                .store(true, std::sync::atomic::Ordering::Relaxed);
1833        }
1834        if self
1835            .graph_failed
1836            .swap(false, std::sync::atomic::Ordering::Relaxed)
1837        {
1838            self.cancel
1839                .store(false, std::sync::atomic::Ordering::Relaxed);
1840            self.clear_sequence_state();
1841            self.graph_logits = None;
1842            self.graph_want_logits = false;
1843            return Err(format!(
1844                "GPU graph failed during NLL {phase} at position {pos}"
1845            ));
1846        }
1847        Ok(())
1848    }
1849
1850    /// Map a virtual layer index to its physical weight index.
1851    /// Looped Transformer (Nanbeige 4.2): 22 physical layers × 2 loops = 44 virtual;
1852    /// virtual layer 23 maps back to physical layer 1 (23 % 22 = 1).
1853    #[inline]
1854    pub fn phys_layer(&self, virtual_idx: usize) -> usize {
1855        virtual_idx % self.physical_layers
1856    }
1857
1858    /// True when `virtual_idx` is the last layer of a loop iteration
1859    /// (used for loop_final_norm insertion).
1860    #[inline]
1861    pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
1862        self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
1863    }
1864
1865    /// Build a pipeline from parts (used by the loader and tests).
1866    #[allow(clippy::too_many_arguments)]
1867
1868    /// Whole-block q1 token graph on the GPU (macOS/Metal): the run of
1869    /// consecutive q1 layers — GDN *and* full attention — starting at
1870    /// `start` executes as few command buffers as the CPU truly needs.
1871    /// Hidden stays device-resident across every layer; the only syncs
1872    /// are before each CPU attend (it needs q/k/v and owns the KV
1873    /// cache) and the final hidden readback. Recurrent states
1874    /// round-trip through shared memory (the CPU stays their owner, so
1875    /// every other path remains coherent). Returns the first layer
1876    /// index NOT covered (== `start` → refused, caller falls through
1877    /// to the per-layer CPU path).
1878    /// Should prefill run position-by-position through the GPU token
1879    /// graph instead of the batched CPU chunk-GEMM? True for q1 GDN
1880    /// hybrids on native Metal: their chunk prefill is walled by the
1881    /// sequential scalar recurrence, so the graph's decode rate wins.
1882    /// NOT for Looped Transformers, despite the per-chunk loop_final_norm
1883    /// sync: the chunk-GEMM amortizes each weight over the whole chunk,
1884    /// which the per-position graph cannot (Nanbeige 4.2 on M4, 512-token
1885    /// prompt: 85 tok/s chunked vs 14 through the graph).
1886    #[cfg(target_os = "macos")]
1887    fn graph_prefill_preferred(&self) -> bool {
1888        let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
1889        if !crate::gpu::enabled_here()
1890            || !graph_force
1891            || std::env::var("CMF_GPU_BLOCK")
1892                .map(|v| v == "0")
1893                .unwrap_or(false)
1894            // CMF_PREFILL_GRAPH=0: the chunked prefill (GEMM projections,
1895            // CPU recurrence) instead of the per-position token graph.
1896            || std::env::var("CMF_PREFILL_GRAPH").as_deref() == Ok("0")
1897        {
1898            return false;
1899        }
1900        self.weights
1901            .layers
1902            .iter()
1903            .any(|lw| {
1904                matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.metal_graph_parts().is_some())
1905            })
1906    }
1907
1908    /// Prompt ingest through the batched wgpu graph in device-prefix mode:
1909    /// a MoE stack that does not fit the card runs each chunk's leading
1910    /// layers on the device (experts resident) and the rest on the host's
1911    /// batched walk. On by default for models with per-layer attention
1912    /// geometry (MiMo-V2 — its measured default); `CMF_BATCH_PREFIX=1`
1913    /// opts any other MoE model in, `=0` keeps the chunked host prefill.
1914    #[cfg(not(target_os = "macos"))]
1915    fn batch_prefix_prefill(&self) -> bool {
1916        let forced = match std::env::var("CMF_BATCH_PREFIX").as_deref() {
1917            Ok("0") => return false,
1918            Ok("1") => true,
1919            _ => false,
1920        };
1921        crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
1922            && crate::gpu::enabled_here()
1923            && !self.graph_refused()
1924            && (forced || self.graph_attn_decline_reason().is_some())
1925            && self.wgpu_graph_attn_decline().is_none()
1926            && self.attn_softcap == 0.0
1927            && self
1928                .weights
1929                .layers
1930                .iter()
1931                .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
1932            && self.automatic_gpu_prefix().is_some()
1933    }
1934
1935    #[cfg(not(target_os = "macos"))]
1936    fn graph_prefill_preferred(&self) -> bool {
1937        // Discrete-GPU wgpu whole-token graph: GDN layers carry recurrent state
1938        // (conv ring + delta-rule S) resident on the GPU. A batched CPU prefill
1939        // builds that state on the CPU only, leaving the GPU buffers zeroed at
1940        // decode → garbage. Route GDN-hybrid prefill through the graph one
1941        // position at a time so the resident state is seeded exactly as decode
1942        // will read it. Pure-attention models keep the batched CPU prefill (its
1943        // KV mirror re-syncs from the CPU cache, so no seeding gap).
1944        let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
1945        if !graph_on || !crate::gpu::enabled_here() {
1946            return false;
1947        }
1948        // Embryo's phase state is device-owned by the resident graph during
1949        // prefill; the batched CPU path would leave decode seeing a zeroed
1950        // device recurrence.  Route the prompt position-by-position too.
1951        if self.embryo_resident_eligible() {
1952            // The resident graph owns the phase state and anchor KV on the
1953            // device.  A prefill-only graph would leave decode on the host
1954            // with no way to import that state, so keep the whole sequence
1955            // on one owner (or use the ordinary CPU prefill/decode pair).
1956            return crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
1957        }
1958        // The descriptor-aware Prism graph now carries both the FWHT/affine
1959        // transforms and resident GDN state, so it is also the exact prefill
1960        // path for this model.  Keeping it here (rather than falling through
1961        // to the CPU chunk walk) is required for a long prompt to seed the
1962        // same device state that decode consumes.
1963        // O(1) needs the CPU prefill: the q-trace that seals the Nyström
1964        // skeleton is recorded there and nowhere else. The GDN half of
1965        // the hybrid loses nothing — the graph's first decode creates
1966        // its (ring, S) entries seeded from `cpu_state`, the same
1967        // handoff every graph run relies on when the entry is fresh.
1968        // Without this line the two designs collide on hybrids and o1
1969        // never becomes graph-portable: prefill through the graph
1970        // records no trace, so views stay None forever.
1971        if self.o1_active() {
1972            return false;
1973        }
1974        // A model the wgpu graphs decline outright would walk its prompt
1975        // one position at a time through a graph that never runs: the
1976        // batched CPU chunk prefill is the right ingest for it.
1977        if self.wgpu_graph_attn_decline().is_some() {
1978            return false;
1979        }
1980        if self
1981            .weights
1982            .layers
1983            .iter()
1984            .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
1985        {
1986            return true;
1987        }
1988        // MoE models too: the chunked CPU prefill runs every expert on the
1989        // host (Hy-MT2-30B-A3B on a Xeon: 8 tok/s of ingest against 53 of
1990        // graph decode), while the token graph — and the batched graph under
1991        // CMF_BATCH_K — keep the experts resident. Full attention in the
1992        // graph writes the KV mirror that decode reads, exactly as it does
1993        // for the hybrids' attention layers. Only when the whole stack is
1994        // resident: with a device prefix the per-position walk finishes
1995        // every token on the host, and the chunked prefill (GEMMs on the
1996        // card, the expert loop batched on the host) is the faster ingest
1997        // (the 8 GB ladder point: 7 tok/s chunked against ~1 walked).
1998        self.weights
1999            .layers
2000            .iter()
2001            .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
2002            && self.automatic_gpu_prefix().is_none()
2003    }
2004
2005    #[cfg(target_os = "macos")]
2006    fn q1_graph_gpu(
2007        &mut self,
2008        start: usize,
2009        upto: Option<usize>,
2010        position: usize,
2011        h: &mut [f32],
2012    ) -> usize {
2013        let _mt0 = std::time::Instant::now(); // CMF_METAL_HOSTPROF
2014        use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
2015        let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
2016        if self.attn_softcap > 0.0 // capped scores: no graph kernel — CPU path
2017            || !crate::gpu::enabled_here()
2018            || !graph_force
2019            || std::env::var("CMF_GPU_BLOCK")
2020                .map(|v| v == "0")
2021                .unwrap_or(false)
2022        {
2023            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2024                eprintln!(
2025                    "block-graph: front gate (softcap={} enabled_here={} graph_force={})",
2026                    self.attn_softcap > 0.0,
2027                    crate::gpu::enabled_here(),
2028                    graph_force,
2029                );
2030            }
2031            if self.graph_head_required {
2032                self.fail_metal_graph("native graph front gate refused");
2033            }
2034            return start;
2035        }
2036        // The graph encodes SiLU or exact-GELU FFNs and attention with an
2037        // explicit model scale. Sliding-window layers with their own RoPE
2038        // table / rotary width ride it too (`metal_graph_swa`): both the
2039        // device attend and the sandwich's CPU attend read each layer's
2040        // window and table. Sandwich norms and other activations fall back
2041        // to the CPU path.
2042        let swa_graph = self.metal_graph_swa();
2043        if (self.swa.is_some() && !swa_graph)
2044            || self.global_attn.is_some()
2045            || self.attention_heads_per_layer.is_some()
2046            || self.attn_v_norm
2047            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
2048            || (self.graph_attn_decline_reason().is_some() && !swa_graph)
2049            || self.weights.layers.iter().any(|lw| {
2050                lw.attn_out_norm.is_some()
2051                    || lw.ffn_out_norm.is_some()
2052                    || lw.layer_scale.is_some()
2053                    || matches!(&lw.ffn, FfnKind::Dense(d) if !matches!(d.act, Act::Silu | Act::Gelu))
2054            })
2055        {
2056            // The Metal graphs' device attend has no per-layer attention
2057            // geometry (the wgpu graphs do): say so once, by name.
2058            if let Some(reason) = self.graph_attn_decline_reason().filter(|_| !swa_graph) {
2059                self.note_graph_decline("metal block graph", reason);
2060            }
2061            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2062                eprintln!(
2063                    "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
2064                    self.swa.is_some(),
2065                    self.global_attn.is_some(),
2066                    self.attention_heads_per_layer.is_some(),
2067                    self.attn_v_norm,
2068                    (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
2069                );
2070            }
2071            if self.graph_head_required {
2072                self.fail_metal_graph("native graph architecture gate refused");
2073            }
2074            return start;
2075        }
2076        // Looped Transformer: the graph covers ALL loop iterations;
2077        // encode_loop_norm is inserted on-device at each boundary.
2078        let limit = upto
2079            .map(|u| u + 1)
2080            .unwrap_or(self.num_layers)
2081            .min(self.num_layers);
2082
2083        enum Item<'a> {
2084            Gdn {
2085                run: Vec<GdnGpuLayer<'a>>,
2086                first: usize,
2087            },
2088            Attn {
2089                l: AttnGpuLayer<'a>,
2090                li: usize,
2091                q_norm: Option<&'a [f32]>,
2092                k_norm: Option<&'a [f32]>,
2093                output_gate: bool,
2094                /// `self_attn.g_proj` output gate (per-head flag): the
2095                /// sandwich projects the device-normed input on the host
2096                /// and gates the attend's output before O.
2097                proj_gate: Option<(&'a QTensor, bool)>,
2098                /// The same gate's f32 rows `[nh × hidden]` when the device
2099                /// attend can apply it (per-head sigmoid, f32 in RAM).
2100                head_gate_w: Option<&'a [f32]>,
2101                bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
2102                /// Attend on the device too (no sync): F32 KV, no
2103                /// o1/bias, dims inside the kernels' contract.
2104                full_gpu: bool,
2105            },
2106        }
2107
2108        // Device-attend KERNEL contract, shared by every Full layer. The
2109        // hd>128 default-off POLICY is applied after the scan: it was
2110        // measured on dense models, and a MoE plan inverts it — with the
2111        // experts on device each CPU-attend sandwich costs a
2112        // commit+wait, ~30 submits/token (W2 on M4: 14.7 tok/s
2113        // sandwiched vs 27.1 device-attend vs 18.8 pure CPU).
2114        let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
2115        let attend_contract = attend_mode != "0"
2116            && attend_mode != "off"
2117            && self.head_dim % 4 == 0
2118            && self.head_dim <= 256
2119            && self.rotary_dim >= 2
2120            && self.rotary_dim <= self.head_dim
2121            && (self.rotary_dim / 2) % 32 == 0
2122            && self.num_kv_heads > 0
2123            && self.num_heads % self.num_kv_heads == 0;
2124
2125        let mut plan: Vec<Item> = Vec::new();
2126        let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
2127        // Break-reason diagnostics ride the same env as the plan summary.
2128        let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
2129        let mut scan = start;
2130        while scan < limit {
2131            let lw = &self.weights.layers[self.phys_layer(scan)];
2132            let ffn = match &lw.ffn {
2133                FfnKind::Dense(d) if d.segs.is_empty() => {
2134                    let (Some(g), Some(u), Some(dn)) = (
2135                        d.gate_proj.metal_graph_parts(),
2136                        d.up_proj.metal_graph_parts(),
2137                        d.down_proj.metal_graph_parts(),
2138                    ) else {
2139                        if block_diag {
2140                            eprintln!(
2141                                "block-graph: L{scan} FFN trio not graph-mappable — run ends"
2142                            );
2143                        }
2144                        break;
2145                    };
2146                    MetalFfn::Dense {
2147                        gate: g,
2148                        up: u,
2149                        down: dn,
2150                        // The front gate admitted SiLU or exact GELU only.
2151                        gelu: d.act == Act::Gelu,
2152                    }
2153                }
2154                FfnKind::Moe(m) => {
2155                    let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
2156                        if block_diag {
2157                            eprintln!(
2158                                "block-graph: L{scan} MoE outside the graph contract — run ends"
2159                            );
2160                        }
2161                        break;
2162                    };
2163                    if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
2164                        model_ref.get_or_insert_with(|| model.clone());
2165                    }
2166                    MetalFfn::Moe(moe)
2167                }
2168                _ => {
2169                    if block_diag {
2170                        eprintln!("block-graph: L{scan} non-graph FFN — run ends");
2171                    }
2172                    break;
2173                }
2174            };
2175            match &lw.attn {
2176                AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
2177                    let parts = (
2178                        w.in_proj_qkv.metal_graph_parts(),
2179                        w.in_proj_z.metal_graph_parts(),
2180                        w.in_proj_a.f32_parts(),
2181                        w.in_proj_b.f32_parts(),
2182                        w.out_proj.metal_graph_parts(),
2183                    );
2184                    let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
2185                        if block_diag {
2186                            eprintln!(
2187                                "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
2188                                w.in_proj_qkv.metal_graph_parts().is_some(),
2189                                w.in_proj_z.metal_graph_parts().is_some(),
2190                                w.in_proj_a.f32_parts().is_some(),
2191                                w.in_proj_b.f32_parts().is_some(),
2192                                w.out_proj.metal_graph_parts().is_some(),
2193                            );
2194                        }
2195                        break;
2196                    };
2197                    if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
2198                        model_ref.get_or_insert_with(|| model.clone());
2199                    }
2200                    let gl = GdnGpuLayer {
2201                        attn_norm: &lw.input_norm,
2202                        post_norm: &lw.post_norm,
2203                        qkv,
2204                        z,
2205                        a,
2206                        b,
2207                        out,
2208                        ffn,
2209                        conv1d: &w.conv1d,
2210                        a_log: &w.a_log,
2211                        dt_bias: &w.dt_bias,
2212                        gnorm: &w.norm,
2213                    };
2214                    match plan.last_mut() {
2215                        Some(Item::Gdn { run, .. }) => run.push(gl),
2216                        _ => plan.push(Item::Gdn {
2217                            run: vec![gl],
2218                            first: scan,
2219                        }),
2220                    }
2221                }
2222                AttnKind::Full {
2223                    wq,
2224                    wk,
2225                    wv,
2226                    wo,
2227                    q_norm,
2228                    k_norm,
2229                    output_gate,
2230                    softplus_gate,
2231                    bias,
2232                } if (!self.kv_cache.layers[scan].o1_sealed()
2233                    // Sealed o1 stays plannable when the Metal o1 port
2234                    // is on: full_gpu attends through the device state,
2235                    // and any refusal falls to the sandwich, whose CPU
2236                    // core routes sealed layers through the nystrom step.
2237                    || std::env::var("CMF_O1_METAL").as_deref() == Ok("1"))
2238                    // A projected output gate: Spark-X2.5's per-head
2239                    // sigmoid form only (Laguna's softplus gate is
2240                    // unmeasured on these paths and keeps the CPU walk).
2241                    && softplus_gate
2242                        .as_ref()
2243                        .is_none_or(|(_, per_head)| *per_head && self.proj_gate_sigmoid) =>
2244                {
2245                    let parts = (
2246                        wq.metal_graph_parts(),
2247                        wk.metal_graph_parts(),
2248                        wv.metal_graph_parts(),
2249                        wo.metal_graph_parts(),
2250                    );
2251                    let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
2252                        break;
2253                    };
2254                    if let QTensor::Mapped { model, .. } = wq {
2255                        model_ref.get_or_insert_with(|| model.clone());
2256                    }
2257                    let cache = &self.kv_cache.layers[scan];
2258                    // O(1) layer on Metal: the device attends through the
2259                    // sealed Nystrom state (opt-in while the port proves
2260                    // itself). Unsealed -> sandwich path = the CPU o1 step.
2261                    let o1_metal = cache.o1.is_some()
2262                        && std::env::var("CMF_O1_METAL").as_deref() == Ok("1")
2263                        && cache.o1_views().is_some();
2264                    // The device attend takes the layer's window and RoPE
2265                    // table, and the per-head gate from f32 rows; a gate
2266                    // held otherwise sandwiches.
2267                    let head_gate_w = softplus_gate
2268                        .as_ref()
2269                        .and_then(|(g, _)| g.f32_parts())
2270                        .filter(|&(_, r, c)| r == self.num_heads && c == self.hidden_size)
2271                        .map(|(d, _, _)| d);
2272                    let full_gpu = attend_contract
2273                        && softplus_gate.is_none() == head_gate_w.is_none()
2274                        && cache.mode == crate::kv_cache::KvMode::F32
2275                        && (cache.o1.is_none() || o1_metal)
2276                        && bias.is_none()
2277                        && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
2278                        && pk.1 == self.num_kv_heads * self.head_dim
2279                        && pv.1 == self.num_kv_heads * self.head_dim
2280                        && po.2 == self.num_heads * self.head_dim;
2281                    plan.push(Item::Attn {
2282                        l: AttnGpuLayer {
2283                            attn_norm: &lw.input_norm,
2284                            post_norm: &lw.post_norm,
2285                            wq: pq,
2286                            wk: pk,
2287                            wv: pv,
2288                            wo: po,
2289                            ffn,
2290                        },
2291                        li: scan,
2292                        q_norm: q_norm.as_deref(),
2293                        k_norm: k_norm.as_deref(),
2294                        output_gate: *output_gate,
2295                        proj_gate: softplus_gate.as_ref().map(|(g, per_head)| (g, *per_head)),
2296                        head_gate_w,
2297                        bias: bias
2298                            .as_ref()
2299                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2300                        full_gpu,
2301                    });
2302                }
2303                _ => break,
2304            }
2305            scan += 1;
2306        }
2307        let Some(model) = model_ref else {
2308            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2309                eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
2310            }
2311            if self.graph_head_required {
2312                self.fail_metal_graph("native graph has no mapped model reference");
2313            }
2314            return start;
2315        };
2316        if plan.is_empty() {
2317            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2318                eprintln!("q1-graph: empty plan at layer {start}");
2319            }
2320            if self.graph_head_required {
2321                self.fail_metal_graph("native graph plan is empty");
2322            }
2323            return start;
2324        }
2325        let has_moe = plan.iter().any(|it| match it {
2326            Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
2327            Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
2328        });
2329        let has_gdn = plan.iter().any(|it| matches!(it, Item::Gdn { .. }));
2330        let dev_attend = attend_contract
2331            && (self.head_dim <= 128
2332                || has_moe
2333                // A GDN hybrid attends on a quarter of its layers: the
2334                // hd>128 caution was measured on pure-dense models where
2335                // gqa_attend dominates, and on Qwen3.8-27B (hd 256, 48
2336                // GDN + 16 attn) the sandwich costs 2x the whole decode
2337                // (1.2 vs 2.21 tok/s measured before the arena fix).
2338                || (self.head_dim <= 256 && has_gdn)
2339                // Sliding-window layers bound most attends by the window:
2340                // Spark-X2.5-1.7B q4tp on the M4 decodes 67 tok/s
2341                // device-attended against 23 sandwiched (28 syncs/token).
2342                || (self.head_dim <= 256 && swa_graph)
2343                || attend_mode == "force"
2344                || attend_mode == "256");
2345        if !dev_attend {
2346            for it in &mut plan {
2347                if let Item::Attn { li, full_gpu, .. } = it {
2348                    // The hd>128 policy is about gqa_attend; an o1 layer
2349                    // attends through its own kernel set.
2350                    let keep_o1 = self.kv_cache.layers[*li].o1.is_some()
2351                        && std::env::var("CMF_O1_METAL").as_deref() == Ok("1");
2352                    if !keep_o1 {
2353                        *full_gpu = false;
2354                    }
2355                }
2356            }
2357        }
2358        if std::env::var("CMF_GRAPH_DBG").is_ok() {
2359            use std::sync::atomic::{AtomicBool, Ordering};
2360            static SAID: AtomicBool = AtomicBool::new(false);
2361            if !SAID.swap(true, Ordering::Relaxed) {
2362                let fg = plan
2363                    .iter()
2364                    .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
2365                    .count();
2366                let att = plan
2367                    .iter()
2368                    .filter(|it| matches!(it, Item::Attn { .. }))
2369                    .count();
2370                eprintln!(
2371                    "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
2372                    plan.len(),
2373                    self.head_dim,
2374                    self.rotary_dim,
2375                    self.num_kv_heads,
2376                    self.num_heads,
2377                );
2378            }
2379        }
2380        let dims = GraphDims {
2381            hidden: self.hidden_size,
2382            eps: self.rms_eps as f32,
2383            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2384        };
2385        let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
2386            if self.graph_head_required {
2387                self.fail_metal_graph("native TokenGraph allocation refused");
2388            }
2389            return start;
2390        };
2391        if swa_graph && self.head_dim > 128 {
2392            // Spark-X2.5 (hd 256, 4 Q heads per KV head): the GQA-shared
2393            // split-K attend reads each K/V row once for the group instead
2394            // of once per head plus the importance re-read — at depth 1000
2395            // on the M4, 59.5 tok/s against 50.0 with the per-head kernel
2396            // up to the default 512 (every sliding layer sits at <= 512).
2397            graph.set_attend_blk_from(64);
2398        }
2399        let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
2400            nv: cfg.num_v_heads,
2401            nk: cfg.num_k_heads,
2402            dk: cfg.key_head_dim,
2403            dv: cfg.value_head_dim,
2404            kk: cfg.conv_kernel,
2405            hidden: self.hidden_size,
2406            inter: self.intermediate_size,
2407            c_dim: cfg.conv_dim(),
2408            eps: cfg.rms_eps as f32,
2409            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2410        });
2411        // Validate the whole plan BEFORE encoding anything: after the
2412        // first sync a refused layer would leave the token
2413        // half-executed, so truncate to the provably encodable prefix.
2414        let mut valid = 0usize;
2415        let mut end = start;
2416        crate::gpu::stageprof(1, _mt0.elapsed()); // конец планирования
2417        if std::env::var("CMF_PLAN_DUMP").is_ok() {
2418            static ONCE: std::sync::Once = std::sync::Once::new();
2419            ONCE.call_once(|| {
2420                for it in &plan {
2421                    match it {
2422                        Item::Gdn { first, run } => {
2423                            eprintln!("plan: Gdn first={first} len={}", run.len())
2424                        }
2425                        Item::Attn { li, full_gpu, .. } => {
2426                            eprintln!("plan: Attn li={li} full_gpu={full_gpu}")
2427                        }
2428                    }
2429                }
2430            });
2431        }
2432        for item in &plan {
2433            let ok = match item {
2434                Item::Gdn { run, .. } => gcfg
2435                    .as_ref()
2436                    .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
2437                    .unwrap_or(false),
2438                Item::Attn { l, .. } => graph.attn_ok(l),
2439            };
2440            if !ok {
2441                if block_diag {
2442                    eprintln!(
2443                        "block-graph: plan item {} ({}) failed graph preflight",
2444                        valid,
2445                        match item {
2446                            Item::Gdn { run, first } => format!("GDN run L{first}+{}", run.len()),
2447                            Item::Attn { li, .. } => format!("Attn L{li}"),
2448                        }
2449                    );
2450                }
2451                break;
2452            }
2453            valid += 1;
2454            end += match item {
2455                Item::Gdn { run, .. } => run.len(),
2456                Item::Attn { .. } => 1,
2457            };
2458        }
2459        plan.truncate(valid);
2460        if plan.is_empty() {
2461            if self.graph_head_required {
2462                self.fail_metal_graph("native graph preflight produced no valid items");
2463            }
2464            return start;
2465        }
2466
2467        if self.graph_head_required && (upto.is_some() || end != self.num_layers) {
2468            self.fail_metal_graph("fused-head NLL requires a complete 64-layer graph");
2469            return start;
2470        }
2471
2472        // Plain dense decode (every item a device-attended full-attention
2473        // layer with a dense FFN, no O(1) state): the only plan shape the
2474        // masked-nibble q4tp matvec and the concurrent layer encoder were
2475        // measured on (MiniCPM5-2B, Qwen3-0.6B on the M4). Hybrids, MoE and
2476        // o1 layers keep the historical serial path bit for bit.
2477        // Every projection must be ONE dispatch (q1t adds an overlay pass,
2478        // Prism q2tp a transform pass — dependent pairs a concurrent
2479        // encoder would race).
2480        let one_pass = |t: (usize, usize, usize)| {
2481            use cortiq_core::TensorDtype as D;
2482            matches!(
2483                model.tensors[t.0].dtype,
2484                D::Q4TiledP | D::Q4Tiled | D::Q4Block | D::Q8Row | D::Q8_2f | D::Q1
2485            )
2486        };
2487        let dense_fast = plan.iter().all(|it| match it {
2488            Item::Attn {
2489                l, li, full_gpu, ..
2490            } => {
2491                *full_gpu
2492                    && self.kv_cache.layers[*li].o1.is_none()
2493                    && [l.wq, l.wk, l.wv, l.wo].into_iter().all(one_pass)
2494                    && match l.ffn {
2495                        MetalFfn::Dense { gate, up, down, .. } => {
2496                            one_pass(gate) && one_pass(up) && one_pass(down)
2497                        }
2498                        _ => false,
2499                    }
2500            }
2501            Item::Gdn { .. } => false,
2502        });
2503        let ab = crate::gpu_metal::dense_ab_arm().filter(|_| dense_fast);
2504        let _mv_fast = match ab {
2505            Some((bits, _)) => {
2506                graph.set_dense_concurrent_raw(bits & crate::gpu_metal::DENSE_CONC != 0);
2507                crate::gpu_metal::MvFastGuard::set_raw(bits)
2508            }
2509            None => {
2510                graph.set_dense_concurrent(dense_fast);
2511                // Spark-X2.5 q8_2f: the four-row q8_2f matvec decodes the
2512                // 1.7B at 40.5 tok/s on the M4 against 27.7 one-row.
2513                let q8r4 = if swa_graph {
2514                    crate::gpu_metal::DENSE_Q8R4
2515                } else {
2516                    0
2517                };
2518                crate::gpu_metal::MvFastGuard::set_bits(if dense_fast {
2519                    crate::gpu_metal::DENSE_MV | crate::gpu_metal::DENSE_FUSE | q8r4
2520                } else {
2521                    0
2522                })
2523            }
2524        };
2525
2526        // RoPE table and rotary width are per layer (`layer_inv_freq`,
2527        // `layer_geom`) inside the loop below.
2528        let pool = self.pool.clone();
2529        let (nh, nkv, hd, hs, eps) = (
2530            self.num_heads,
2531            self.num_kv_heads,
2532            self.head_dim,
2533            self.hidden_size,
2534            self.rms_eps,
2535        );
2536        let norm_style = self.norm_style;
2537        let gemma = norm_style == cortiq_core::NormStyle::Gemma;
2538        let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
2539        let kv_id = self.graph_kv_id;
2540        // GDN runs whose states await readback after the next sync
2541        // (device-attended layers add no sync, so several may stack).
2542        let mut pending: Vec<(usize, usize)> = Vec::new();
2543        // Device-attended layers: their K/V/imp are pulled from the
2544        // mirror after the final sync.
2545        let mut dev_attn: Vec<usize> = Vec::new();
2546        for item in &plan {
2547            let _xt0 = std::time::Instant::now();
2548            let _xkind: u32 = match item {
2549                Item::Gdn { .. } => 2,
2550                Item::Attn { .. } => 3,
2551            };
2552            // Looped Transformer: insert on-device norm at loop boundaries.
2553            if self.loop_final_norm {
2554                let item_start = match item {
2555                    Item::Gdn { first, .. } => *first,
2556                    Item::Attn { li, .. } => *li,
2557                };
2558                if item_start > start && self.is_loop_end(item_start - 1) {
2559                    graph.encode_loop_norm(&self.weights.final_norm);
2560                }
2561            }
2562            match item {
2563                Item::Gdn { run, first } => {
2564                    for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
2565                        if l.linear_state.len() != want {
2566                            l.linear_state = vec![0f32; want];
2567                        }
2568                    }
2569                    let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
2570                        .iter()
2571                        .map(|l| l.linear_state.as_slice())
2572                        .collect();
2573                    let _ig = std::time::Instant::now();
2574                    if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
2575                        // Unreachable: the plan was validated above.
2576                        tracing::error!("q1 graph: GDN run refused after validation");
2577                        return start;
2578                    }
2579                    // Early commit: the GPU starts the run while the
2580                    // CPU encodes the next layer (nothing to wait on).
2581                    graph.commit_kind = 2;
2582                    graph.commit();
2583                    crate::gpu::stageprof(0, _ig.elapsed());
2584                    pending.push((*first, run.len()));
2585                }
2586                Item::Attn {
2587                    l,
2588                    li,
2589                    q_norm,
2590                    k_norm,
2591                    output_gate,
2592                    proj_gate,
2593                    head_gate_w,
2594                    bias,
2595                    full_gpu,
2596                } => {
2597                    let _ia = std::time::Instant::now();
2598                    // This layer's attention geometry: its window, RoPE
2599                    // table and rotary width (Spark-X2.5 interleaves
2600                    // 512-window layers rotating all dims at θ 1e4 with
2601                    // full layers rotating a quarter at θ 5e6). Every model
2602                    // without sliding layers reads the global ones here.
2603                    let inv_freq_l = self.layer_inv_freq(*li);
2604                    let rd_l = self.layer_geom(*li).2;
2605                    let window_l = self.layer_window(*li);
2606                    let head_gate_w = *head_gate_w;
2607                    // ── Fully device-resident attention: no sync at all.
2608                    if *full_gpu {
2609                        let cache = &self.kv_cache.layers[*li];
2610                        let o1p = if cache.o1.is_some() {
2611                            match cache.o1_views() {
2612                                Some(views) => Some(crate::gpu::O1AttnParams {
2613                                    views,
2614                                    epoch: self.o1_epoch,
2615                                }),
2616                                // Sealed state gone mid-run: sandwich.
2617                                None => None,
2618                            }
2619                        } else {
2620                            None
2621                        };
2622                        let o1_layer = cache.o1.is_some();
2623                        if o1_layer && o1p.is_none() {
2624                            // fall to the sandwich (CPU o1 step)
2625                        }
2626                        let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
2627                        let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
2628                        let cpu_stored = if o1_layer { 0 } else { cpu_k[0].len() / hd };
2629                        let p = crate::gpu::AttnDeviceParams {
2630                            kv_id,
2631                            layer: *li,
2632                            nh,
2633                            nkv,
2634                            hd,
2635                            rd: rd_l,
2636                            position,
2637                            scale: self.attn_scale,
2638                            eps: eps as f32,
2639                            gemma,
2640                            late_qk_norm: self.qk_norm_after_rope,
2641                            output_gate: *output_gate,
2642                            q_norm: *q_norm,
2643                            k_norm: *k_norm,
2644                            inv_freq: &inv_freq_l,
2645                            cpu_k,
2646                            cpu_v,
2647                            cpu_stored,
2648                            o1: o1p,
2649                            window: window_l,
2650                            head_gate: head_gate_w,
2651                        };
2652                        let o1_bad = o1_layer && p.o1.is_none();
2653                        if !o1_bad && graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p)
2654                        {
2655                            // o1 layers leave no mirror row to pull.
2656                            if p.o1.is_none() {
2657                                dev_attn.push(*li);
2658                            }
2659                            graph.commit_kind = 3;
2660                            graph.commit();
2661                            // The footer below is skipped by `continue`:
2662                            // account the device-attn item here or its
2663                            // cost hides from the stage profile entirely.
2664                            crate::gpu::stageprof(_xkind, _xt0.elapsed());
2665                            continue;
2666                        }
2667                        // Mirror refused (nothing encoded) → sandwich.
2668                    }
2669                    graph.encode_attn_prefix(l);
2670                    if let Err(err) = graph.sync_checked() {
2671                        self.fail_metal_graph(&err);
2672                        return start;
2673                    }
2674                    if !pending.is_empty() {
2675                        let idxs: Vec<usize> =
2676                            pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2677                        let mut outs: Vec<&mut [f32]> = self
2678                            .kv_cache
2679                            .layers
2680                            .iter_mut()
2681                            .enumerate()
2682                            .filter(|(i, _)| idxs.binary_search(i).is_ok())
2683                            .map(|(_, s)| s.linear_state.as_mut_slice())
2684                            .collect();
2685                        graph.read_states(&mut outs);
2686                    }
2687                    let mut q_raw = attention::take_buf(l.wq.1);
2688                    let mut k = attention::take_buf(l.wk.1);
2689                    let mut v = attention::take_buf(l.wv.1);
2690                    graph.read_qkv(&mut q_raw, &mut k, &mut v);
2691                    // Projected output gate: g = G·norm(h) on the host, off
2692                    // the normed input the prefix left on the device.
2693                    let mut gate_raw = proj_gate.map(|(gp, _)| {
2694                        let mut normed = attention::take_buf(hs);
2695                        graph.read_normed(&mut normed);
2696                        let mut raw = attention::take_buf(gp.rows());
2697                        gp.matvec(&normed, &mut raw, pool.as_deref());
2698                        attention::recycle_buf(&mut normed);
2699                        raw
2700                    });
2701                    let cfg = QwenAttnCfg {
2702                        num_heads: nh,
2703                        num_kv_heads: nkv,
2704                        head_dim: hd,
2705                        hidden_size: hs,
2706                        position,
2707                        inv_freq: &inv_freq_l,
2708                        rotary_dim: rd_l,
2709                        scale: self.attn_scale,
2710                        softcap: self.attn_softcap,
2711                        window: window_l,
2712                        v_norm: false,
2713                        qk_norm_after_rope: self.qk_norm_after_rope,
2714                        gate_sigmoid: self.proj_gate_sigmoid,
2715                        q_norm: *q_norm,
2716                        k_norm: *k_norm,
2717                        output_gate: *output_gate,
2718                        softplus_gate: None,
2719                        rope_scale: 1.0,
2720                        bias: *bias,
2721                        rms_eps: eps,
2722                        norm_style,
2723                        pool: pool.as_deref(),
2724                        v_head_dim: hd,
2725                    };
2726                    // CMF_ATTN_ORACLE=1: diff the device attend against
2727                    // this CPU attend on identical inputs (bring-up).
2728                    let oracle = std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")
2729                        || std::env::var("CMF_ATTN_DUMP").is_ok();
2730                    let _ = full_gpu;
2731                    let oracle_in = oracle.then(|| (q_raw.clone(), k.clone(), v.clone()));
2732                    let mut ao = attention::qwen_attention_core(
2733                        q_raw,
2734                        k,
2735                        v,
2736                        &mut self.kv_cache.layers[*li],
2737                        &cfg,
2738                    );
2739                    // CMF_ATTN_DUMP=<dir>: this token's rope'd Q and the layer's whole
2740                    // K/V cache as raw f32 (offline attention-statistics probes:
2741                    // block bounds, mass concentration). Needs CMF_GPU_ATTEND=0.
2742                    if let Ok(dir) = std::env::var("CMF_ATTN_DUMP") {
2743                        if let Some((qr0, k0, v0)) = oracle_in.clone() {
2744                            let (cq, _cg, _ck, _cv) =
2745                                attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2746                            let cache = &self.kv_cache.layers[*li];
2747                            let n = cache.head_keys(0).len() / hd;
2748                            let mut bytes: Vec<u8> = Vec::new();
2749                            for v in [nh as u32, nkv as u32, hd as u32, n as u32, position as u32] {
2750                                bytes.extend_from_slice(&v.to_le_bytes());
2751                            }
2752                            for v in &cq {
2753                                bytes.extend_from_slice(&v.to_le_bytes());
2754                            }
2755                            for g in 0..nkv {
2756                                for v in cache.head_keys(g) {
2757                                    bytes.extend_from_slice(&v.to_le_bytes());
2758                                }
2759                            }
2760                            for g in 0..nkv {
2761                                for v in cache.head_values(g) {
2762                                    bytes.extend_from_slice(&v.to_le_bytes());
2763                                }
2764                            }
2765                            let _ =
2766                                std::fs::write(format!("{dir}/L{li}_pos{position}.bin"), &bytes);
2767                        }
2768                    }
2769                    if let Some((qr0, k0, v0)) =
2770                        oracle_in.filter(|_| std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1"))
2771                    {
2772                        let (cq, _cg, ck, cv) =
2773                            attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2774                        let mut h_now = vec![0f32; hs];
2775                        graph.read_h(&mut h_now);
2776                        let cache = &self.kv_cache.layers[*li];
2777                        let n_after = cache.head_keys(0).len() / hd;
2778                        // A sealed O(1) cache may have no dense current-row
2779                        // entry. The oracle is a debug probe, so let it see
2780                        // zero stored exact rows instead of underflowing.
2781                        let stored = n_after.saturating_sub(1);
2782                        let cpu_k: Vec<&[f32]> = (0..nkv)
2783                            .map(|g| &cache.head_keys(g)[..stored * hd])
2784                            .collect();
2785                        let cpu_v: Vec<&[f32]> = (0..nkv)
2786                            .map(|g| &cache.head_values(g)[..stored * hd])
2787                            .collect();
2788                        let p = crate::gpu::AttnDeviceParams {
2789                            kv_id,
2790                            layer: *li,
2791                            nh,
2792                            nkv,
2793                            hd,
2794                            // The layer's own geometry, as the CPU attend
2795                            // above used it (the projected gate is applied
2796                            // after this probe on both sides).
2797                            rd: rd_l,
2798                            position,
2799                            scale: self.attn_scale,
2800                            eps: eps as f32,
2801                            gemma,
2802                            late_qk_norm: self.qk_norm_after_rope,
2803                            output_gate: *output_gate,
2804                            q_norm: *q_norm,
2805                            k_norm: *k_norm,
2806                            inv_freq: &inv_freq_l,
2807                            cpu_k,
2808                            cpu_v,
2809                            cpu_stored: stored,
2810                            o1: None,
2811                            window: window_l,
2812                            head_gate: None,
2813                        };
2814                        if let Some((dq, dk, dv, dao)) = graph.debug_attn_device(l, &p, &h_now) {
2815                            let md = |a: &[f32], b: &[f32]| {
2816                                a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs()))
2817                            };
2818                            let nn = |a: &[f32]| a.iter().map(|x| x * x).sum::<f32>().sqrt();
2819                            eprintln!(
2820                                "attn-oracle L{li} pos {position}: |q| {:.2} max|dq| {:.4} | |k| {:.2} max|dk| {:.4} | |v| {:.2} max|dv| {:.4} | |ao| {:.2} max|dao| {:.4}",
2821                                nn(&cq),
2822                                md(&cq, &dq),
2823                                nn(&ck),
2824                                md(&ck, &dk),
2825                                nn(&cv),
2826                                md(&cv, &dv),
2827                                nn(&ao),
2828                                md(&ao, &dao)
2829                            );
2830                        } else {
2831                            eprintln!("attn-oracle L{li}: device probe declined");
2832                        }
2833                    }
2834                    if let (Some(raw), Some((_, per_head))) = (gate_raw.as_deref(), *proj_gate) {
2835                        // V is as wide as the head (the front gate), so ao
2836                        // is nh·hd here, as in `qwen_attention`.
2837                        attention::apply_projected_gate(
2838                            &mut ao,
2839                            raw,
2840                            per_head,
2841                            hd,
2842                            self.proj_gate_sigmoid,
2843                        );
2844                    }
2845                    if let Some(mut raw) = gate_raw.take() {
2846                        attention::recycle_buf(&mut raw);
2847                    }
2848                    graph.encode_attn_suffix(l, &ao);
2849                    // Early commit: the GPU starts O+FFN while the CPU
2850                    // encodes the following GDN run / attention prefix.
2851                    graph.commit();
2852                    attention::recycle_buf(&mut ao);
2853                }
2854            }
2855
2856            crate::gpu::stageprof(_xkind, _xt0.elapsed());
2857        }
2858        // Ride the final norm + lm_head in the same command buffer when
2859        // this run reaches the model's end and the caller wants logits:
2860        // the separate per-op lm_head submit (a full round trip) folds
2861        // into the sync that already happens here.
2862        let mut lm_rows = None;
2863        if self.graph_want_logits
2864            && upto.is_none()
2865            && end == self.num_layers
2866            && std::env::var("CMF_GPU_LMHEAD")
2867                .map(|v| v != "0")
2868                .unwrap_or(true)
2869        {
2870            if let Some(lm) = self.weights.lm_head.metal_graph_parts() {
2871                if graph.lm_head_ok(lm) {
2872                    graph.encode_lm_head(&self.weights.final_norm, lm);
2873                    lm_rows = Some(lm.1);
2874                }
2875            }
2876        }
2877        if self.graph_head_required && lm_rows.is_none() {
2878            METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2879            self.fail_metal_graph("fused graph head was requested but not encodable");
2880            return start;
2881        }
2882        let _sy0 = std::time::Instant::now();
2883        if let Err(err) = graph.sync_checked() {
2884            self.fail_metal_graph(&err);
2885            return start;
2886        }
2887        let _rs0 = std::time::Instant::now();
2888        if !pending.is_empty() {
2889            let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2890            let mut outs: Vec<&mut [f32]> = self
2891                .kv_cache
2892                .layers
2893                .iter_mut()
2894                .enumerate()
2895                .filter(|(i, _)| idxs.binary_search(i).is_ok())
2896                .map(|(_, s)| s.linear_state.as_mut_slice())
2897                .collect();
2898            graph.read_states(&mut outs);
2899        }
2900        if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() == Ok("1") {
2901            use std::sync::atomic::{AtomicU64, Ordering};
2902            static SY: AtomicU64 = AtomicU64::new(0);
2903            static RS: AtomicU64 = AtomicU64::new(0);
2904            static N: AtomicU64 = AtomicU64::new(0);
2905            SY.fetch_add((_rs0 - _sy0).as_nanos() as u64, Ordering::Relaxed);
2906            RS.fetch_add(_rs0.elapsed().as_nanos() as u64, Ordering::Relaxed);
2907            let n = N.fetch_add(1, Ordering::Relaxed) + 1;
2908            if n % 100 == 0 {
2909                eprintln!(
2910                    "postprof: sync-wait {:.1} ms/ток | read_states {:.1} ms/ток ({n})",
2911                    SY.load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
2912                    RS.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
2913                );
2914            }
2915        }
2916        if let Some(rows) = lm_rows {
2917            crate::gpu::hostprof_encode_done(_mt0);
2918            let mut lg = attention::take_buf(rows.min(self.vocab_size));
2919            graph.read_logits(&mut lg);
2920            crate::gpu::hostprof_total(_mt0);
2921            lg.resize(self.vocab_size, 0.0);
2922            if let Some(c) = self.final_softcap {
2923                for l in lg.iter_mut() {
2924                    *l = c * (*l / c).tanh();
2925                }
2926            }
2927            self.graph_logits = Some(lg);
2928        }
2929        graph.read_h(h);
2930        if self.graph_head_required && self.graph_logits.is_none() {
2931            METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2932            self.fail_metal_graph("fused graph head completed without logits readback");
2933            return start;
2934        }
2935        METAL_GRAPH_TOK_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2936        METAL_GRAPH_LAYERS.fetch_add(
2937            end.saturating_sub(start) as u64,
2938            std::sync::atomic::Ordering::Relaxed,
2939        );
2940        if self.graph_head_required {
2941            METAL_GRAPH_HEAD_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2942        }
2943        // Device-attended layers: replay the CPU bookkeeping — append
2944        // the mirror's new K/V row (rope'd on the GPU) into the owner
2945        // cache, then bank this token's attention-importance mass.
2946        for li in dev_attn {
2947            let mut krow = attention::take_buf(nkv * hd);
2948            let mut vrow = attention::take_buf(nkv * hd);
2949            if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
2950                let cache = &mut self.kv_cache.layers[li];
2951                cache.append(&krow, &vrow, &[]);
2952                let n = cache.seq_len;
2953                let mut imp = attention::take_buf(n);
2954                crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
2955                cache.accumulate_imp(&imp);
2956                attention::recycle_buf(&mut imp);
2957            }
2958            attention::recycle_buf(&mut krow);
2959            attention::recycle_buf(&mut vrow);
2960        }
2961        if let Some((_, arm)) = ab {
2962            crate::gpu_metal::dense_ab_record(arm, _mt0.elapsed());
2963        }
2964        end
2965    }
2966
2967    pub fn new(
2968        tokenizer: Tokenizer,
2969        weights: PipelineWeights,
2970        hidden_size: usize,
2971        intermediate_size: usize,
2972        num_heads: usize,
2973        num_kv_heads: usize,
2974        head_dim: usize,
2975        num_layers: usize,
2976        physical_layers: usize,
2977        loop_final_norm: bool,
2978        vocab_size: usize,
2979        rms_eps: f64,
2980        rope_base: f32,
2981        norm_style: NormStyle,
2982        max_seq_len: usize,
2983        sampler_config: SamplerConfig,
2984    ) -> Self {
2985        let rng = match sampler_config.seed {
2986            Some(s) => SplitMix64::new(s),
2987            None => SplitMix64::from_entropy(),
2988        };
2989        let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
2990        let pool = Pool::from_env();
2991        if let Some(p) = &pool {
2992            tracing::info!("worker pool: {} threads", p.n_workers());
2993            // Keep the workers on the socket that holds the weights.
2994            if let Some(model) = weights
2995                .lm_head
2996                .model_arc()
2997                .or_else(|| weights.embed_tokens.model_arc())
2998            {
2999                let regions: Vec<&[u8]> =
3000                    model.tensors.iter().map(|t| model.entry_bytes(t)).collect();
3001                p.bind_numa(&regions);
3002            }
3003        }
3004        Self {
3005            gpu_plan: None,
3006            tokenizer: std::sync::Arc::new(tokenizer),
3007            kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
3008            sampler_config,
3009            weights,
3010            hidden_size,
3011            intermediate_size,
3012            num_heads,
3013            num_kv_heads,
3014            head_dim,
3015            num_layers,
3016            physical_layers,
3017            loop_final_norm,
3018            vocab_size,
3019            rms_eps,
3020            rope_base,
3021            norm_style,
3022            rotary_dim: head_dim,
3023            attention_heads_per_layer: None,
3024            kv_heads_per_layer: None,
3025            v_head_dim: None,
3026            layer_dump: std::env::var_os("CMF_LAYER_DUMP")
3027                .filter(|v| !v.is_empty())
3028                .map(std::path::PathBuf::from),
3029            graph_declines: std::cell::RefCell::new(Vec::new()),
3030            mimo_moe: Default::default(),
3031            vmf_cfg: None,
3032            gdn_cfg: None,
3033            kda_cfg: None,
3034            g3n: None,
3035            dsv4: None,
3036            dsv41: None,
3037            dsv41_vision: None,
3038            dsv41_prefill: None,
3039            qwen4_exp: None,
3040            dsv4_mtp: Vec::new(),
3041            dspark: None,
3042            dspark_pending: Vec::new(),
3043            dspark_hist: Vec::new(),
3044            dspark_real: Vec::new(),
3045            dspark_trunk_picks: Vec::new(),
3046            dspark_exp: Vec::new(),
3047            dspark_draft_ns: 0,
3048            logit_multiplier: None,
3049            cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
3050            graph_failed: std::sync::atomic::AtomicBool::new(false),
3051            kv_history: Vec::new(),
3052            kv_history_device: false,
3053            short_conv_cfg: None,
3054            mtp: None,
3055            mimo_mtp: None,
3056            verify_exact_moe: false,
3057            speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
3058            ignore_eos: false,
3059            draft_full_streak: 0,
3060            spec_k_adapt: None,
3061            spec_acc_ewma: 0.7,
3062            rng,
3063            sampler_scratch: SamplerScratch::default(),
3064            spec_forced: None,
3065            spec_q: Vec::new(),
3066            spec_p: Vec::new(),
3067            spec_res: Vec::new(),
3068            spec_qs: Vec::new(),
3069            spec_ps: Vec::new(),
3070            spec_ress: Vec::new(),
3071            mtp_graph_mode: None,
3072            #[cfg(target_os = "macos")]
3073            metal_verify: None,
3074            inv_freq,
3075            ws: ForwardScratch::new(hidden_size),
3076            pool,
3077            model: None,
3078            dyn_force_f32: false,
3079            dyn_skill_layers: Vec::new(),
3080            dyn_active: None,
3081            dyn_blend_loaded: false,
3082            dyn_phi_layer: None,
3083            dyn_phi_ema: Vec::new(),
3084            dyn_phi_seen: 0,
3085            dyn_router: None,
3086            o1_cfg: None,
3087            o1_epoch: 0,
3088            o1_flags: Vec::new(),
3089            trace: false,
3090            calib_temp: 1.0,
3091            confidence_on: true,
3092            embed_multiplier: 1.0,
3093            attn_scale: 1.0 / (head_dim as f32).sqrt(),
3094            swa: None,
3095            sliding_layers: None,
3096            anchor_core: None,
3097            bounded_rope: None,
3098            kv_prefix: KvPrefix::default(),
3099            last_prefill_tokens: 0,
3100            inv_freq_local: None,
3101            rotary_dim_local: None,
3102            rope_scale: 1.0,
3103            rope_scale_local: 1.0,
3104            global_attn: None,
3105            inv_freq_global: None,
3106            attn_v_norm: false,
3107            qk_norm_after_rope: false,
3108            proj_gate_sigmoid: false,
3109            final_softcap: None,
3110            head_clusters: None,
3111            attn_softcap: 0.0,
3112            graph_want_logits: false,
3113            graph_head_required: false,
3114            graph_logits: None,
3115            embryo_graph: None,
3116            graph_refused: std::sync::atomic::AtomicBool::new(false),
3117            graph_kv_id: {
3118                static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
3119                NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
3120            },
3121            #[cfg(test)]
3122            nll_test_fail_at: None,
3123            #[cfg(test)]
3124            nll_test_force_serial: false,
3125        }
3126    }
3127
3128    /// Enable/disable per-layer O(1) Nyström attention. Only Full
3129    /// layers are eligible (a linear layer keeps its own operator).
3130    /// Applies to generation (`generate*`/`forward_ids`): the prompt
3131    /// pass stays exact, then the state seals after prefill or at the
3132    /// deferred skeleton-safe boundary for short prompts; decode runs on
3133    /// the O(1) state. Teacher-forced scoring (`ppl_ids`) intentionally
3134    /// stays exact.
3135    pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
3136        if let Err(e) = self.try_set_o1(cfg) {
3137            tracing::error!("{e}");
3138        }
3139    }
3140
3141    /// True when the file carries a native bounded anchor
3142    /// (`arch.anchor_core`): its state is a fixed record the header
3143    /// fixes, and the post-hoc O(1) override is meaningless on it.
3144    pub fn bounded_native(&self) -> bool {
3145        self.anchor_core.is_some()
3146    }
3147
3148    /// Bytes the resident device graph holds for this pipeline's sequence:
3149    /// `(recurrent state, anchor KV/ring)`; None on the host path.
3150    pub fn device_state_bytes(&self) -> Option<(u64, u64)> {
3151        crate::gpu::embryo_device_state_bytes(self.graph_kv_id)
3152    }
3153
3154    /// Why an O(1) override is refused on this pipeline, if it is.
3155    pub fn o1_refusal(&self) -> Option<String> {
3156        self.anchor_core.as_ref().map(|ac| {
3157            format!(
3158                "--o1 / CMF_O1 refused: the anchor is native bounded \
3159                 (anchor_core kind={} window={} sink={}); the file's operator \
3160                 is executed as-is and no post-hoc Nyström overlay applies",
3161                ac.kind, ac.window, ac.sink
3162            )
3163        })
3164    }
3165
3166    /// `set_o1` that reports the refusal instead of logging it.
3167    pub fn try_set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) -> Result<(), String> {
3168        if let Some(c) = &cfg {
3169            if let Some(why) = self.o1_refusal() {
3170                self.o1_flags = Vec::new();
3171                self.o1_cfg = None;
3172                return Err(why);
3173            }
3174            if crate::nystrom::o1_deferred_boundary(c.w, c.sink).is_none() {
3175                self.o1_flags.clear();
3176                self.o1_cfg = None;
3177                return Err(format!(
3178                    "o1 disabled: w + sink + slack + 1 overflows usize (w={}, sink={})",
3179                    c.w, c.sink
3180                ));
3181            }
3182        }
3183        self.o1_flags = match &cfg {
3184            Some(c) => {
3185                let mut flags = c.layer_flags(self.num_layers);
3186                for (li, f) in flags.iter_mut().enumerate() {
3187                    // The Nyström state replaces a full-context plain
3188                    // softmax: a sliding window or a learned sink is not
3189                    // something it can represent, and a V narrower than
3190                    // the head is not what its streaming state stores.
3191                    // Those layers keep exact cache attention.
3192                    if *f
3193                        && (!matches!(
3194                            self.weights.layers[self.phys_layer(li)].attn,
3195                            AttnKind::Full { .. }
3196                        ) || self.layer_window(li).is_some()
3197                            || self.kv_cache.layers[li].sinks.is_some()
3198                            || self.layer_v_dim(li) != self.layer_geom(li).1)
3199                    {
3200                        *f = false;
3201                    }
3202                }
3203                flags
3204            }
3205            None => Vec::new(),
3206        };
3207        if let Some(c) = &cfg {
3208            let n = self.o1_flags.iter().filter(|&&f| f).count();
3209            tracing::info!(
3210                "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
3211                self.num_layers,
3212                c.m,
3213                c.w,
3214                c.sink,
3215                c.rect
3216            );
3217        }
3218        self.o1_cfg = cfg;
3219        Ok(())
3220    }
3221
3222    /// Install the file's bounded anchor: one fixed-size ring per
3223    /// `AttnKind::Bounded` layer (from the header, not per prompt) and
3224    /// the shared relative-rotation table. Must run after the RoPE setup
3225    /// (`set_rotary`, YaRN) so the table is built from the final
3226    /// `inv_freq`.
3227    pub fn install_bounded(
3228        &mut self,
3229        cfg: &cortiq_core::AnchorCoreConfig,
3230    ) -> Result<(), String> {
3231        if !cortiq_core::AnchorCoreConfig::KINDS.contains(&cfg.kind.as_str()) {
3232            return Err(format!(
3233                "anchor_core kind '{}' is not executable by this runtime",
3234                cfg.kind
3235            ));
3236        }
3237        if cfg.window == 0 {
3238            return Err("anchor_core.window must be >= 1".into());
3239        }
3240        let mut n = 0usize;
3241        for li in 0..self.num_layers {
3242            let pl = self.phys_layer(li);
3243            if let AttnKind::Bounded(w) = &self.weights.layers[pl].attn {
3244                if w.window != cfg.window || w.sink != cfg.sink {
3245                    return Err(format!(
3246                        "layer {li}: bounded weights (window {} sink {}) disagree with \
3247                         anchor_core (window {} sink {})",
3248                        w.window, w.sink, cfg.window, cfg.sink
3249                    ));
3250                }
3251                self.kv_cache.layers[li].install_bounded(cfg.window);
3252                n += 1;
3253            }
3254        }
3255        if n == 0 {
3256            return Err("anchor_core is present but no layer executes it".into());
3257        }
3258        let rope = crate::bounded::BoundedRope::new(cfg.window, &self.inv_freq, self.rope_scale);
3259        self.bounded_rope = Some(std::sync::Arc::new(rope));
3260        self.anchor_core = Some(cfg.clone());
3261        self.embryo_graph = None;
3262        tracing::info!(
3263            "bounded anchor {}: {n} layer(s), window {} sink {} — {} B of ring per layer",
3264            cfg.kind,
3265            cfg.window,
3266            cfg.sink,
3267            self.kv_cache.layers.iter().map(|l| l.bounded_state_bytes()).max().unwrap_or(0)
3268        );
3269        Ok(())
3270    }
3271
3272    /// Tag every layer's cache with its wire record kind and the model's
3273    /// operator identity (hash64 of `linear_core_identity` JSON) so the
3274    /// versioned state wire refuses a peer holding another operator.
3275    pub fn install_wire_identity(&mut self, identity: u64) {
3276        for li in 0..self.kv_cache.layers.len() {
3277            let pl = self.phys_layer(li);
3278            let kind = match self.weights.layers.get(pl).map(|l| &l.attn) {
3279                Some(AttnKind::Bounded(_)) => crate::kv_cache::WireKind::Bounded,
3280                Some(AttnKind::Linear(_))
3281                | Some(AttnKind::LinearGdn(_))
3282                | Some(AttnKind::ShortConv(_))
3283                | Some(AttnKind::Kda(_)) => crate::kv_cache::WireKind::Linear,
3284                _ => crate::kv_cache::WireKind::Full,
3285            };
3286            let l = &mut self.kv_cache.layers[li];
3287            l.wire_kind = kind;
3288            l.wire_identity = identity;
3289            // A per-layer geometry (`set_attn_geometry`, Gemma's global
3290            // heads) rebuilds a layer cache with index 0: re-tag it.
3291            l.wire_layer = li as u32;
3292        }
3293    }
3294
3295    /// Forget the reuse keys (legacy `kv_history` and the bounded prefix).
3296    pub fn clear_history(&mut self) {
3297        self.kv_history.clear();
3298        self.kv_history_device = false;
3299        self.kv_prefix.clear();
3300    }
3301
3302    /// Did THIS pipeline's token graph refuse for a structural reason?
3303    /// (Per pipeline: another lane's refusal, or a new pipeline of the
3304    /// same model, never changes it.)
3305    pub fn graph_refused(&self) -> bool {
3306        self.graph_refused
3307            .load(std::sync::atomic::Ordering::Relaxed)
3308    }
3309
3310    /// Remember a structural refusal of this pipeline's token graph.
3311    pub fn mark_graph_refused(&self) {
3312        if !self
3313            .graph_refused
3314            .swap(true, std::sync::atomic::Ordering::Relaxed)
3315        {
3316            tracing::info!(
3317                "token graph: unsupported for this pipeline (seq {}) — not retrying",
3318                self.graph_kv_id
3319            );
3320        }
3321    }
3322
3323    /// Position the resident Embryo graph holds for this pipeline's
3324    /// sequence (`Some(next position)`), `None` when the device holds no
3325    /// image of it (host-owned sequence, or none started).
3326    pub fn device_sequence_position(&self) -> Option<usize> {
3327        crate::gpu::embryo_device_next_position(self.graph_kv_id)
3328    }
3329
3330    /// Is the sequence this pipeline's reuse key describes owned by the
3331    /// device path it would take now? A prefix recorded on the resident
3332    /// graph continues only there (at exactly `n`), a host prefix only on
3333    /// the host; any mismatch means the next turn re-prefills from zero.
3334    fn prefix_owner_matches(&self, n: usize, recorded_on_device: bool) -> bool {
3335        let dev = self.device_sequence_position();
3336        if recorded_on_device {
3337            dev == Some(n) && self.embryo_resident_wanted()
3338        } else {
3339            dev.is_none()
3340        }
3341    }
3342
3343    /// The weights under the sequence changed (a real skill switch): every
3344    /// cached state was computed by other weights. Clear the host KV /
3345    /// ring / recurrent state AND the reuse keys — a surviving `kv_prefix`
3346    /// would let the next turn "extend" a prefix the new weights never
3347    /// saw — drop the packed resident graph (it holds the old FFN
3348    /// tensors; the next build packs the live ones under a fresh id) and
3349    /// reset its device sequence.
3350    pub(crate) fn invalidate_for_weight_change(&mut self) {
3351        self.clear_sequence_state();
3352        self.embryo_graph = None;
3353    }
3354
3355    /// Prompt positions the cache already holds when `input_ids`
3356    /// strictly EXTENDS the consumed prefix (0 otherwise). Bounded-native
3357    /// models answer from the fixed-size `kv_prefix` record; everything
3358    /// else from the legacy `kv_history` vector (which the network split
3359    /// also reads and writes).
3360    fn cached_prefix_len(&self, input_ids: &[u32]) -> usize {
3361        let (n, on_device) = if self.bounded_native() {
3362            (self.kv_prefix.extension(input_ids), self.kv_prefix.on_device())
3363        } else {
3364            let h = &self.kv_history;
3365            if !h.is_empty() && h.len() < input_ids.len() && input_ids[..h.len()] == h[..] {
3366                (h.len(), self.kv_history_device)
3367            } else {
3368                (0, false)
3369            }
3370        };
3371        // The owner tag: the state the key describes must live where this
3372        // turn will continue it (R4) — else a fresh sequence.
3373        if n > 0 && !self.prefix_owner_matches(n, on_device) {
3374            tracing::warn!(
3375                "kv-reuse refused: the cached prefix ({n} positions) was built on the {} path, \
3376                 the device now holds {:?} — re-prefilling from zero",
3377                if on_device { "resident device" } else { "host" },
3378                self.device_sequence_position()
3379            );
3380            return 0;
3381        }
3382        n
3383    }
3384
3385    /// Public view of the prefix-reuse decision for `input_ids` (positions
3386    /// the next `generate*` would take from the cache; 0 = fresh sequence).
3387    pub fn reusable_prefix_len(&self, input_ids: &[u32]) -> usize {
3388        self.cached_prefix_len(input_ids)
3389    }
3390
3391    /// Record the forwarded prefix as the next turn's reuse key. A
3392    /// bounded-native model extends the rolling record (`reused` = the
3393    /// positions this turn found cached); others keep the literal vector.
3394    /// Either way the key carries its OWNER: the resident device graph
3395    /// (it holds an image of this sequence) or the host.
3396    fn record_consumed_prefix(&mut self, consumed: &[u32], reused: usize) {
3397        let on_device = self.device_sequence_position().is_some();
3398        if self.bounded_native() {
3399            let keep = reused > 0 && reused == self.kv_prefix.len() && reused <= consumed.len();
3400            let prev_device = self.kv_prefix.on_device();
3401            self.kv_history.clear();
3402            self.kv_history_device = false;
3403            if keep && prev_device == on_device {
3404                self.kv_prefix.extend(&consumed[reused..]);
3405            } else {
3406                self.kv_prefix.set(consumed);
3407            }
3408            self.kv_prefix.set_on_device(on_device);
3409        } else {
3410            self.kv_history = consumed.to_vec();
3411            self.kv_history_device = on_device;
3412        }
3413    }
3414
3415    /// True when at least one layer runs the O(1) kernel.
3416    pub fn o1_active(&self) -> bool {
3417        self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
3418    }
3419
3420    /// Whether generation's prompt ingest is routed through the whole-token
3421    /// graph.  The bench uses this to label the measured generation prefill
3422    /// honestly; keep the predicate in Pipeline so CLI labels cannot drift
3423    /// from the production route.
3424    /// Positions per batched-graph submit for the prompt: `CMF_BATCH_K`
3425    /// when set (0 = one position at a time through the token graph),
3426    /// otherwise 32 on a discrete card whose prompt takes the graph route.
3427    /// The batched graph read a 2048-token prompt at 53 tok/s against 28.5
3428    /// one position at a time on an RTX PRO 4000 (Qwen3.8-27B q4tp: TTFT
3429    /// 39 s against 72), and its states are the speculative verify's,
3430    /// measured identical to the plain path. macOS keeps its own arm.
3431    pub fn generation_batch_k(&self) -> usize {
3432        if let Some(k) = std::env::var("CMF_BATCH_K")
3433            .ok()
3434            .and_then(|v| v.parse::<usize>().ok())
3435        {
3436            return k;
3437        }
3438        #[cfg(not(target_os = "macos"))]
3439        if self.graph_prefill_preferred() && !self.o1_active() {
3440            return 32;
3441        }
3442        0
3443    }
3444
3445    pub fn generation_graph_prefill(&self) -> bool {
3446        let graph = self.graph_prefill_preferred();
3447        // On wgpu, an active MTP head now consumes the trunk's graph batches
3448        // and warms its own block from those returned rows.  The selected
3449        // generation measurement is therefore the batched path, even though
3450        // the underlying GDN model still satisfies the graph-prefill
3451        // predicate.  Keep the CLI label tied to the actual route.  Native
3452        // Metal has a separate prefill-batch arm and retains its historical
3453        // label here.
3454        // A batched prompt (`generation_batch_k` > 0) is the batched graph
3455        // for every model on the graph route, not only those with an MTP
3456        // head — the label follows the route.
3457        #[cfg(not(target_os = "macos"))]
3458        if graph
3459            && self.generation_batch_k() > 0
3460            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
3461        {
3462            return false;
3463        }
3464        graph
3465    }
3466
3467    /// Device-side O(1) mirrors currently uploaded for this pipeline's
3468    /// sequence.  The count/bytes are zero before seal or after a fresh
3469    /// reset; callers use this to distinguish logical host state from the
3470    /// GPU allocation that actually serves decode.
3471    pub fn o1_device_stats(&self) -> (usize, u64) {
3472        crate::gpu::o1_device_stats(self.graph_kv_id)
3473    }
3474
3475    /// Arm query collection on the o1 layers (fresh prompt pass).
3476    /// Reset the o1 layers to Collecting for a fresh sequence. Pub for the
3477    /// network split: each side runs the o1 lifecycle over ITS OWN layers
3478    /// (begin before prefill, seal at the prefill barrier).
3479    pub fn o1_begin(&mut self) {
3480        self.o1_begin_with_prefix(None);
3481    }
3482
3483    /// Arm collection and optionally request a positive calibration prefix.
3484    /// The effective barrier is always at least the skeleton-safe floor, so
3485    /// a short requested prefix cannot create an exact-only runtime state.
3486    pub fn o1_begin_with_prefix(&mut self, requested_prefix: Option<usize>) {
3487        if let Some(c) = &self.o1_cfg {
3488            let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
3489            let boundary = requested_prefix.map(|p| {
3490                p.max(
3491                    crate::nystrom::o1_deferred_boundary(w, sink)
3492                        .expect("o1 config boundary validated in set_o1"),
3493                )
3494            });
3495            for (li, &f) in self.o1_flags.iter().enumerate() {
3496                if f {
3497                    self.kv_cache.layers[li].o1_begin_with_boundary(m, w, sink, rect, boundary);
3498                }
3499            }
3500        }
3501    }
3502
3503    /// Effective deferred boundary for a positive prefix request.
3504    fn o1_effective_boundary(&self, requested_prefix: usize) -> Option<usize> {
3505        self.o1_cfg.as_ref().and_then(|c| {
3506            crate::nystrom::o1_deferred_boundary(c.w, c.sink)
3507                .map(|floor| requested_prefix.max(floor))
3508        })
3509    }
3510
3511    fn o1_note_transition(&mut self) {
3512        // Drain every layer's one-shot bit before publishing one pipeline
3513        // epoch. `any()` would short-circuit on the first layer and leak the
3514        // remaining bits into later forwards, causing one epoch per layer.
3515        let mut transitioned = false;
3516        for (li, &flagged) in self.o1_flags.iter().enumerate() {
3517            if flagged {
3518                transitioned |= self.kv_cache.layers[li].take_o1_transition();
3519            }
3520        }
3521        if transitioned {
3522            self.o1_epoch = self.o1_epoch.wrapping_add(1);
3523        }
3524    }
3525
3526    fn o1_pending(&self) -> bool {
3527        self.o1_flags.iter().enumerate().any(|(li, &f)| {
3528            f && self.kv_cache.layers[li].seq_len > 0
3529                && self.kv_cache.layers[li].o1_pending_boundary().is_some()
3530        })
3531    }
3532
3533    fn o1_fail(&mut self, err: String) {
3534        tracing::error!("o1 deferred seal failed; terminating sequence: {err}");
3535        self.clear_sequence_state();
3536        self.graph_failed
3537            .store(true, std::sync::atomic::Ordering::Relaxed);
3538        self.cancel
3539            .store(true, std::sync::atomic::Ordering::Relaxed);
3540    }
3541
3542    /// Seal participating layers while retaining the exact state when the
3543    /// prompt is below the deferred boundary. A split worker may have
3544    /// collecting layers outside its owned span; zero-depth layers remain
3545    /// armed and are intentionally skipped until their peer runs them.
3546    pub fn o1_seal_checked(&mut self) -> Result<bool, String> {
3547        if self.o1_cfg.is_none() {
3548            return Ok(false);
3549        }
3550        let mut participating = false;
3551        for li in 0..self.num_layers {
3552            if !self.o1_flags.get(li).copied().unwrap_or(false) {
3553                continue;
3554            }
3555            if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3556                return Err(err);
3557            }
3558            if self.kv_cache.layers[li].seq_len == 0 {
3559                continue;
3560            }
3561            participating = true;
3562            let num_heads = self.layer_num_heads(li);
3563            self.kv_cache.layers[li].o1_seal_checked(num_heads)?;
3564        }
3565        self.o1_note_transition();
3566        for li in 0..self.num_layers {
3567            if self.o1_flags.get(li).copied().unwrap_or(false) {
3568                if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3569                    return Err(err);
3570                }
3571            }
3572        }
3573        Ok(participating
3574            && (0..self.num_layers).all(|li| {
3575                !self.o1_flags.get(li).copied().unwrap_or(false)
3576                    || self.kv_cache.layers[li].seq_len == 0
3577                    || self.kv_cache.layers[li].o1_sealed()
3578            }))
3579    }
3580
3581    /// Complete a deferred boundary after a full position/span forward.
3582    /// This is the pipeline owner for epoch publication and failure cleanup.
3583    fn o1_progress(&mut self) {
3584        if !self.o1_active() {
3585            return;
3586        }
3587        for li in 0..self.num_layers {
3588            if self.o1_flags.get(li).copied().unwrap_or(false) {
3589                if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3590                    self.o1_fail(err);
3591                    return;
3592                }
3593            }
3594        }
3595        // A qwen_attention row can seal in the middle of a complete layer
3596        // walk. Consume its transition even though the pending boundary has
3597        // already disappeared from the cache.
3598        self.o1_note_transition();
3599        if !self.o1_pending() {
3600            return;
3601        }
3602        if let Err(err) = self.o1_seal_checked() {
3603            self.o1_fail(err);
3604        }
3605    }
3606
3607    /// Turn a deferred O(1) failure raised by a hidden-only forward into the
3608    /// Result error its public batch/span caller must return. The failure
3609    /// path already cleared host/device sequence state; consume only the
3610    /// side-channel marker here and leave the pipeline reusable.
3611    fn check_o1_progress_failure(&mut self, phase: &str) -> Result<(), String> {
3612        if self
3613            .graph_failed
3614            .swap(false, std::sync::atomic::Ordering::Relaxed)
3615        {
3616            self.cancel
3617                .store(false, std::sync::atomic::Ordering::Relaxed);
3618            self.clear_sequence_state();
3619            return Err(format!("{phase}: deferred O(1) transition failed"));
3620        }
3621        Ok(())
3622    }
3623
3624    /// Freeze landmarks + skeleton state after the prompt pass and drop
3625    /// the o1 layers' full KV; decode then runs `step()` per token.
3626    /// Pub for the network split (see `o1_begin`).
3627    pub fn o1_seal(&mut self) {
3628        if let Err(err) = self.o1_seal_checked() {
3629            self.o1_fail(err);
3630        }
3631    }
3632
3633    /// Enable/disable the structured per-token telemetry trace (B4).
3634    pub fn set_trace(&mut self, on: bool) {
3635        self.trace = on;
3636    }
3637
3638    /// Replace all request-scoped sampler options and reset the random stream.
3639    /// This is required for deterministic `seed` semantics in pooled servers.
3640    pub fn set_sampler_config(&mut self, config: SamplerConfig) {
3641        self.rng = match config.seed {
3642            Some(seed) => SplitMix64::new(seed),
3643            None => SplitMix64::from_entropy(),
3644        };
3645        self.sampler_config = config;
3646    }
3647
3648    /// Toggle the per-token confidence reduction (a full-vocab
3649    /// softmax each token). `bench --core` turns it off so the timed
3650    /// loop matches llama-bench's core contract; the result's
3651    /// `confidence` vec is empty while off.
3652    pub fn set_confidence(&mut self, on: bool) {
3653        self.confidence_on = on;
3654    }
3655
3656    /// Set the confidence-calibration temperature (B1). Values ≤0 are
3657    /// clamped to raw (1.0).
3658    pub fn set_calib_temp(&mut self, t: f32) {
3659        self.calib_temp = if t > 1e-3 { t } else { 1.0 };
3660    }
3661
3662    /// The active calibration temperature (1.0 = raw probability).
3663    pub fn calib_temp(&self) -> f32 {
3664        self.calib_temp
3665    }
3666
3667    /// Partial rotary (Qwen3.5): rotate only the first `rotary_dim` dims;
3668    /// the frequency table is rebuilt over the rotary dims.
3669    pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
3670        self.rotary_dim = rotary_dim.min(self.head_dim);
3671        self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
3672        // The packed resident graph owns its own inverse-frequency plane;
3673        // changing RoPE after it was built must not leave a stale device
3674        // model behind the exact host configuration.
3675        self.embryo_graph = None;
3676    }
3677
3678    fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
3679        QwenAttnCfg {
3680            num_heads: self.num_heads,
3681            num_kv_heads: self.num_kv_heads,
3682            head_dim: self.head_dim,
3683            hidden_size: self.hidden_size,
3684            position,
3685            inv_freq: &self.inv_freq,
3686            rotary_dim: self.rotary_dim,
3687            scale: self.attn_scale,
3688            softcap: self.attn_softcap,
3689            window: None,
3690            v_norm: false,
3691            qk_norm_after_rope: self.qk_norm_after_rope,
3692            gate_sigmoid: self.proj_gate_sigmoid,
3693            q_norm: None,
3694            k_norm: None,
3695            output_gate: false,
3696            softplus_gate: None,
3697            rope_scale: self.rope_scale,
3698            bias: None,
3699            rms_eps: self.rms_eps,
3700            norm_style: self.norm_style,
3701            pool: self.pool.as_deref(),
3702            v_head_dim: self.v_head_dim.unwrap_or(self.head_dim),
3703        }
3704    }
3705
3706    /// Generate text from a plain-text prompt. Streams tokens via `on_token`.
3707    pub fn generate(
3708        &mut self,
3709        prompt: &str,
3710        max_tokens: usize,
3711        task_mask: Option<&TaskMask>,
3712        on_token: Option<TokenCallback>,
3713    ) -> Result<GenerateResult, String> {
3714        let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
3715        self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
3716    }
3717
3718    /// Generate from a V4.1 multimodal prompt prepared by the vision module.
3719    /// Vision rows are encoded once and fed through the same bounded token walk as text.
3720    pub fn generate_from_vl(
3721        &mut self,
3722        input: &crate::dsv41_vision::PreparedVlInputs,
3723        max_tokens: usize,
3724        task_mask: Option<&TaskMask>,
3725        on_token: Option<TokenCallback>,
3726    ) -> Result<GenerateResult, String> {
3727        let Some(dsv41) = &self.dsv41 else {
3728            return Err("V4.1 multimodal input requires a DeepSeek-V4.1 pipeline".into());
3729        };
3730        if input.token_ids.is_empty() {
3731            return Err("empty V4.1 multimodal prompt".into());
3732        }
3733        if input.token_types.len() != input.token_ids.len() {
3734            return Err(format!(
3735                "V4.1 token type count {} != token count {}",
3736                input.token_types.len(),
3737                input.token_ids.len()
3738            ));
3739        }
3740        let dim = dsv41.2.dim;
3741        let mut embeddings = vec![None; input.token_ids.len()];
3742        let mut participates = vec![true; input.token_ids.len()];
3743        if !input.images.is_empty() {
3744            let vision = self
3745                .dsv41_vision
3746                .as_ref()
3747                .ok_or_else(|| "V4.1 image prompt has no loaded vision tower".to_string())?;
3748            for image in &input.images {
3749                let end = image.start.saturating_add(image.types.len());
3750                if end > input.token_ids.len() {
3751                    return Err(format!(
3752                        "V4.1 image span {}..{} exceeds prompt length {}",
3753                        image.start,
3754                        end,
3755                        input.token_ids.len()
3756                    ));
3757                }
3758                let mut span = vec![0.0f32; image.types.len() * dim];
3759                vision.fill_image_span(image, &mut span, self.pool.as_deref())?;
3760                for (offset, &kind) in image.types.iter().enumerate() {
3761                    let pos = image.start + offset;
3762                    if input.token_types[pos] != kind {
3763                        return Err(format!(
3764                            "V4.1 image type mismatch at position {pos}: {} != {kind}",
3765                            input.token_types[pos]
3766                        ));
3767                    }
3768                    embeddings[pos] = Some(span[offset * dim..(offset + 1) * dim].to_vec());
3769                    participates[pos] = false;
3770                }
3771            }
3772        }
3773        for (pos, &kind) in input.token_types.iter().enumerate() {
3774            if kind == crate::dsv41_vision::TEXT && embeddings[pos].is_some() {
3775                return Err(format!("V4.1 text position {pos} has an image embedding"));
3776            }
3777            if kind != crate::dsv41_vision::TEXT && embeddings[pos].is_none() {
3778                return Err(format!("V4.1 image position {pos} has no image embedding"));
3779            }
3780        }
3781        self.dsv41_prefill = Some((embeddings, participates));
3782        let result = self.generate_from_ids(&input.token_ids, max_tokens, task_mask, on_token);
3783        self.dsv41_prefill = None;
3784        result
3785    }
3786
3787    /// `None` when the mask forbids nothing (see `TaskMask::fully_open`).
3788    fn drop_open_mask<'m>(&self, m: Option<&'m TaskMask>) -> Option<&'m TaskMask> {
3789        m.filter(|m| !m.fully_open(self.intermediate_size, self.num_heads))
3790    }
3791
3792    /// Generate from prepared token ids (e.g. a chat template).
3793    ///
3794    /// With an MTP head, greedy generation without a task mask takes the
3795    /// speculative path: the MTP module drafts the token after next and
3796    /// the main model verifies both in one fused two-position forward
3797    /// (weights streamed once). The output is EXACTLY the vanilla greedy
3798    /// sequence — a rejected draft is rolled back — MTP only buys speed.
3799    pub fn generate_from_ids(
3800        &mut self,
3801        input_ids: &[u32],
3802        max_tokens: usize,
3803        task_mask: Option<&TaskMask>,
3804        on_token: Option<TokenCallback>,
3805    ) -> Result<GenerateResult, String> {
3806        self.generate_with_prompt_rows(input_ids, None, max_tokens, task_mask, on_token)
3807    }
3808
3809    /// Generate from complete prompt embeddings [token_count, hidden_size].
3810    /// Text rows can be obtained with `embed_id`; media rows replace only
3811    /// their expanded placeholder positions. Rows are already scaled and
3812    /// enter `PrefillIn::Hidden`, so a device graph must not re-embed them.
3813    /// Token-only KV reuse is disabled both into and out of this request.
3814    pub fn generate_from_embeds(
3815        &mut self,
3816        input_ids: &[u32],
3817        prompt_rows: &[f32],
3818        max_tokens: usize,
3819        task_mask: Option<&TaskMask>,
3820        on_token: Option<TokenCallback>,
3821    ) -> Result<GenerateResult, String> {
3822        if input_ids.is_empty()
3823            || input_ids.len().checked_mul(self.hidden_size) != Some(prompt_rows.len())
3824        {
3825            return Err("embedded prompt dimensions must be [tokens, hidden_size]".into());
3826        }
3827        if prompt_rows.iter().any(|x| !x.is_finite()) {
3828            return Err("embedded prompt contains non-finite values".into());
3829        }
3830        if !self.can_prefill_batched() || self.dyn_router.is_some()
3831            || self.o1_active() || self.mtp.is_some() || self.gpu_plan.is_some()
3832        {
3833            return Err("embedded prompts require the ordinary transformer path without O(1), dynamic routing, GPU splitting or a generic MTP head".into());
3834        }
3835        self.generate_with_prompt_rows(input_ids, Some(prompt_rows), max_tokens, task_mask, on_token)
3836    }
3837
3838    fn generate_with_prompt_rows(
3839        &mut self,
3840        input_ids: &[u32],
3841        prompt_rows: Option<&[f32]>,
3842        max_tokens: usize,
3843        task_mask: Option<&TaskMask>,
3844        mut on_token: Option<TokenCallback>,
3845    ) -> Result<GenerateResult, String> {
3846        #[cfg(target_os = "macos")]
3847        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
3848        if std::env::var("CMF_TRACE_H").is_ok() {
3849            eprintln!("input_ids: {input_ids:?}");
3850        }
3851        if input_ids.is_empty() {
3852            return Err("empty prompt: nothing to generate from".to_string());
3853        }
3854        // A prior graph failure is terminal for that sequence but must not
3855        // poison the next independent request.  Keep this flag separate from
3856        // the externally-owned cooperative cancel bit.
3857        self.graph_failed
3858            .store(false, std::sync::atomic::Ordering::Relaxed);
3859        // A mask that forbids nothing still costs every fused path and
3860        // whole-token graph, all of which are gated on `is_none()`. A
3861        // narrowed file whose one segment is always on carries exactly
3862        // such a mask — drop it here rather than pay 5x for a no-op.
3863        let task_mask = self.drop_open_mask(task_mask);
3864
3865        // Cross-turn KV reuse: a chat app resends the whole history
3866        // every turn; when the new ids strictly EXTEND what the cache
3867        // already holds, prefill only the tail — turn latency stays
3868        // proportional to the new text instead of the whole session.
3869        // Extension-only (no rollback), so it is exact for every layer
3870        // kind including recurrent state; MTP/o1/task-mask runs keep
3871        // the fresh-sequence path. CMF_KV_REUSE=0 disables.
3872        let mut reuse_from = {
3873            let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
3874            if on
3875                && prompt_rows.is_none()
3876                && task_mask.is_none()
3877                && self.mtp.is_none()
3878                && !(self.mimo_mtp.is_some() && self.speculative)
3879                && self.o1_cfg.is_none()
3880                && self.dsv41.is_none()
3881            {
3882                self.cached_prefix_len(input_ids)
3883            } else {
3884                0
3885            }
3886        };
3887        // The device may own rows the host tail prefill needs (wgpu decode
3888        // writes only its mirror): hand them to the host, or start fresh.
3889        if reuse_from > 0 && !self.prepare_kv_reuse(reuse_from) {
3890            reuse_from = 0;
3891        }
3892        self.last_prefill_tokens = input_ids.len() - reuse_from;
3893        let bounded_native = self.bounded_native();
3894        if reuse_from == 0 {
3895            // Fresh sequence — the cache holds absolute positions.
3896            self.clear_sequence_state();
3897        } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
3898            eprintln!(
3899                "kv-reuse: {} of {} prompt positions already cached",
3900                reuse_from,
3901                input_ids.len()
3902            );
3903        }
3904        crate::gpu::graph_race_begin_generation();
3905        // Optional bounded calibration prefix. Keep the requested value
3906        // even when it is longer than the prompt; the collecting layer will
3907        // defer at the effective boundary and remain exact for short input.
3908        let o1_prefill = if self.o1_active() && task_mask.is_none() {
3909            std::env::var("CMF_O1_PREFILL")
3910                .ok()
3911                .and_then(|v| v.parse::<usize>().ok())
3912                .filter(|&p| p > 0)
3913        } else {
3914            None
3915        };
3916        if task_mask.is_none() {
3917            self.o1_begin_with_prefix(o1_prefill);
3918        }
3919
3920        // Speculative decode is off under o1: a rejected draft can't be
3921        // rolled back out of the far accumulators / ring window (the
3922        // Nyström insertion is irreversible by design).
3923        // The wgpu token graph owns a device K/V mirror that speculative
3924        // rollback would desync — the two are mutually exclusive.
3925        let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
3926        // Graph speculative decode (`CMF_GRAPH_SPEC=1`): the MTP head
3927        // drafts, ONE batched graph submit verifies the whole chain.
3928        //
3929        // It now PAYS on Qwen3.6-27B / RTX 5090 — 51.1 tok/s against a
3930        // plain 49.4 at k=3, medians of three, 89% of drafts accepted,
3931        // and the greedy continuation is byte-identical to the plain
3932        // path. That took the batch matvec sharing its nibble unpack
3933        // across the batch (`CMF_MV_BK=2`); before it, the same round
3934        // measured 43.6, an 11% LOSS, which is what the earlier note
3935        // here described.
3936        //
3937        // Still opt-in. One model's win is not a default: the verify
3938        // rides `gdn_spec_restore` and a batched frame whose numerics
3939        // are the batch kernels', and that has to be shown on more than
3940        // one architecture before every greedy decode takes it.
3941        // Greedy (with or without penalties) verifies by argmax equality.
3942        // Sampling (temperature > 0) can go through speculative SAMPLING —
3943        // draft from the MTP head's own post-chain distribution, accept
3944        // with min(1, p/q), correct from max(0, p − q); the emitted stream
3945        // is distributed exactly as the plain sampler's — but it is
3946        // OPT-IN (`CMF_GRAPH_SPEC_SAMPLE=1`): measured on Qwen3.8-27B /
3947        // RTX 5090 at the instruct row (0.7 / 0.80 / 20 / presence 1.5)
3948        // it decoded 19-22 tok/s against a plain 40 — nine post-chain
3949        // distributions a round plus a lower acceptance than greedy's,
3950        // against a verify that costs 2.7 single tokens. The greedy arms
3951        // pay +10%; the sampling arm needs a cheaper verify first.
3952        // Native Metal HAS that verify: its eight-row tile is flat in b,
3953        // so a round costs ~1.9 plain tokens and the sampling arm pays at
3954        // 2.3 accepted per round — measured on Qwen3.8-27B q4tp / M4 at
3955        // the CLI defaults (0.7 / rep 1.1 / top-k 40, seed 42), a code
3956        // prompt: 9.0 tok/s against a plain 5.4 in the same window, and
3957        // the per-round watchdog turns it off where prose loses. So on
3958        // Metal the sampling arm is ON (`CMF_GRAPH_SPEC_SAMPLE=0` opts out)
3959        // — but only for a config the SPARSE chain serves (a top-k within
3960        // `sparse_ok`): without it a round builds nine 248k-float
3961        // distributions on the host, which is the 5090's measured loss and
3962        // not a cost the round-token proxy below can see. A top-k-less
3963        // sampling config keeps the plain path unless asked for by name.
3964        #[cfg(target_os = "macos")]
3965        let metal_graph = crate::gpu::q1_force()
3966            && crate::gpu::enabled_here()
3967            && std::env::var("CMF_GPU_BLOCK")
3968                .map(|v| v != "0")
3969                .unwrap_or(true);
3970        #[cfg(not(target_os = "macos"))]
3971        let metal_graph = false;
3972        let spec_sample_env = std::env::var("CMF_GRAPH_SPEC_SAMPLE").ok();
3973        // A round whose cost is the MEASURED one: greedy (argmax rows), or
3974        // sampling through the sparse chain. Anything else pays the dense
3975        // chain's host time, which no proxy can price.
3976        let spec_cheap_round = self.sampler_config.temperature < 1e-6
3977            || sampler::sparse_ok(&self.sampler_config);
3978        let spec_sampling_ok = self.sampler_config.temperature < 1e-6
3979            || match spec_sample_env.as_deref() {
3980                Some("1") => true,
3981                Some(_) => false,
3982                None => metal_graph && spec_cheap_round,
3983            };
3984        // ON by default for greedy on the wgpu graph: with the draft on
3985        // the graph and the verify bit-exact, it measured 58.7 tok/s
3986        // against a plain 48.1 on Qwen3.8-27B q4tp / RTX 5090 (k=4) and
3987        // 51.1 against 49.4 on Qwen3.6-27B, and a round that stops
3988        // paying turns itself off below (acceptance watchdog).
3989        // `CMF_GRAPH_SPEC=0` disables; `=1` was the old opt-in spelling.
3990        // …but only where the batched verify has its register-blocked
3991        // kernel: q4tp dense FFNs (graph kind 6). q4t and q8_2f verify
3992        // through tile GEMMs today and measured a LOSS (q8_2f 22 against
3993        // 29 tok/s), the 2-bit plane the same; those stay opt-in
3994        // (`CMF_GRAPH_SPEC=1`).
3995        // …at least in nine dense FFNs of ten: a healed file carries its
3996        // last two layers at q8_2f, and two tile-GEMM verifies among 64 do
3997        // not change the arithmetic (measured: the healed q4tp file
3998        // decodes at the plain file's rate and would otherwise sit out).
3999        let (mut dense_n, mut dense_q4tp) = (0usize, 0usize);
4000        for lw in &self.weights.layers {
4001            if let FfnKind::Dense(d) = &lw.ffn {
4002                dense_n += 1;
4003                if matches!(d.gate_proj.graph_weight(), Some((_, _, 6, _)))
4004                    && matches!(d.up_proj.graph_weight(), Some((_, _, 6, _)))
4005                    && matches!(d.down_proj.graph_weight(), Some((_, _, 6, _)))
4006                {
4007                    dense_q4tp += 1;
4008                }
4009            }
4010        }
4011        let spec_default_ok = dense_n == 0 || dense_q4tp * 10 >= dense_n * 9;
4012        // Penalties break the draft head's agreement with the trunk (a
4013        // 1.1 repetition penalty measured 2 of 16 accepted): not by
4014        // default there either — off Metal that rule is untouched, and
4015        // suppressed ids keep counting as a penalty there, because no
4016        // measurement on a discrete card says otherwise.
4017        //
4018        // On native Metal the penalized arms DO pay: the draft applies
4019        // the same penalty and the verify scores the penalized rows
4020        // exactly (`greedy_pen`, the plain loop's arithmetic), so the
4021        // text is the plain path's and only the round's shape changes.
4022        // Measured on this M4 — see the report for the interleaved run.
4023        let penalized = !metal_graph
4024            && (self.sampler_config.repetition_penalty != 1.0
4025                || self.sampler_config.presence_penalty != 0.0
4026                || !self.sampler_config.suppress_tokens.is_empty());
4027        // …and not on wgpu-over-Metal: the batched verify graph there
4028        // returned 0 accepted drafts and garbage text on a GDN hybrid
4029        // (16.08, Qwen3.5-0.8B) while Vulkan is bit-exact; the Mac's
4030        // default backend is native Metal without a batch graph anyway.
4031        #[cfg(feature = "gpu")]
4032        let metal_wgpu = graph_on && crate::gpu_wgpu::wgpu_backend_is_metal();
4033        #[cfg(not(feature = "gpu"))]
4034        let metal_wgpu = false;
4035        let spec_env = std::env::var("CMF_GRAPH_SPEC").ok();
4036        let spec_wanted = match spec_env.as_deref() {
4037            Some("0") => false,
4038            Some(_) => {
4039                if metal_wgpu {
4040                    tracing::warn!(
4041                        "CMF_GRAPH_SPEC forced on wgpu/Metal: the batched verify graph is not \
4042                         verified on this backend (garbage measured on Qwen3.5-0.8B)"
4043                    );
4044                }
4045                true
4046            }
4047            None => spec_default_ok && !penalized && !metal_wgpu,
4048        };
4049        // Native Metal: the b-row verify graph (`try_batch_graph_metal`)
4050        // stands where the wgpu batch graph stands on discrete cards
4051        // (`metal_graph`, above).
4052        let graph_spec = self.speculative
4053            && (graph_on || metal_graph)
4054            && self.mtp.is_some()
4055            && task_mask.is_none()
4056            && !self.o1_active()
4057            && spec_sampling_ok
4058            && spec_wanted;
4059        // Native Metal: say the route ONCE (RUST_LOG=info), so a user can
4060        // confirm the fast path without setting a single flag — every
4061        // knob below defaults to the measured-best value on the M4.
4062        #[cfg(target_os = "macos")]
4063        if metal_graph {
4064            static SAID: std::sync::Once = std::sync::Once::new();
4065            SAID.call_once(|| {
4066                let spec = if graph_spec {
4067                    let k = std::env::var("CMF_GRAPH_SPEC_K")
4068                        .ok()
4069                        .and_then(|v| v.parse::<usize>().ok())
4070                        .filter(|&v| (1..=8).contains(&v))
4071                        .unwrap_or(7);
4072                    let arm = if self.sampler_config.temperature < 1e-6 {
4073                        "greedy"
4074                    } else {
4075                        "sampling"
4076                    };
4077                    format!(
4078                        "spec k={k} {arm} (batched verify, draft shortlist {}, trial: proxy)",
4079                        Self::draft_vocab_rows(usize::MAX)
4080                    )
4081                } else if !self.speculative {
4082                    "spec off (CMF_MTP=0)".to_string()
4083                } else if self.mtp.is_none() {
4084                    "spec off (no MTP head)".to_string()
4085                } else if !spec_sampling_ok {
4086                    if spec_cheap_round {
4087                        "spec off (CMF_GRAPH_SPEC_SAMPLE=0)".to_string()
4088                    } else {
4089                        "spec off (sampling without a top-k: the dense chain \
4090                         costs more than it saves)"
4091                            .to_string()
4092                    }
4093                } else if !spec_wanted {
4094                    "spec off (CMF_GRAPH_SPEC=0 or non-q4tp FFNs)".to_string()
4095                } else if task_mask.is_some() {
4096                    "spec off (task mask)".to_string()
4097                } else {
4098                    "spec off (O(1) attention)".to_string()
4099                };
4100                let on = |var: &str| {
4101                    if std::env::var(var).as_deref() == Ok("0") {
4102                        "off"
4103                    } else {
4104                        "on"
4105                    }
4106                };
4107                tracing::info!(
4108                    "metal native: {spec}, state4 {}, async replay {}, prefill graph {}, \
4109                     MTP graph {}, attend {}, probe {}",
4110                    if crate::gpu_metal::state4_on() { "on" } else { "off" },
4111                    if crate::gpu_metal::async_replay_on() { "on" } else { "off" },
4112                    on("CMF_METAL_PREFILL"),
4113                    on("CMF_MTP_GRAPH"),
4114                    std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into()),
4115                    if crate::gpu::probe_enabled() { "bypassed (q1 force)" } else { "off" },
4116                );
4117            });
4118        }
4119        // GDN hybrids sit the fused-pair speculation out by default: the
4120        // recurrence is sequential, so the pair lane cannot parallelize
4121        // (the bench's own Pair line reads fused 1.28x TWO singles on the
4122        // 35B) and the draft's full-vocab head rides on top — measured 2x
4123        // SLOWER end to end (16.1 vs 32.4 tok/s on the 48-core stand).
4124        // CMF_MTP=1 forces it back for study.
4125        let pair_pays = self.gdn_cfg.is_none() || std::env::var("CMF_MTP").as_deref() == Ok("1");
4126        let spec_active = self.speculative
4127            && self.mtp.is_some()
4128            && task_mask.is_none()
4129            && !self.o1_active()
4130            && ((!graph_on && pair_pays && self.sampler_config.temperature < 1e-6) || graph_spec);
4131        // The MTP module is detached during generation so its mutable
4132        // state does not fight the borrow on `self`.
4133        let mut mtp = if spec_active { self.mtp.take() } else { None };
4134        if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
4135            eprintln!(
4136                "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
4137                mtp.is_some(),
4138                self.speculative,
4139                self.sampler_config.temperature < 1e-6,
4140            );
4141        }
4142        if let Some(m) = &mut mtp {
4143            m.kv.clear();
4144            // The MTP block's own device mirror starts over with its cache.
4145            crate::gpu::graph_kv_reset(self.mtp_kv_id());
4146            self.mtp_graph_mode = None;
4147        }
4148        // MiMo-V2's draft stack: greedy rounds (draft K with the chained
4149        // MTP layers, verify K+1 rows in one batched forward). Sampling
4150        // decodes plain; `CMF_MTP=0` / `CMF_MIMO_MTP=0` turn it off.
4151        let mimo_spec = self.speculative
4152            && self.mimo_mtp.is_some()
4153            && task_mask.is_none()
4154            && !self.o1_active()
4155            && self.dyn_router.is_none()
4156            && self.sampler_config.temperature < 1e-6
4157            && std::env::var("CMF_MIMO_MTP").as_deref() != Ok("0");
4158        if let Some(st) = self.mimo_mtp.as_mut() {
4159            st.reset();
4160            if mimo_spec && std::env::var_os("CMF_MIMO_MTP_PROBE").is_some() {
4161                Self::mimo_mtp_hist_cap(st, input_ids.len());
4162            }
4163        }
4164        // Dynamic router detached during decode (same borrow trick as MTP).
4165        // Speculative decode and dynamic routing are mutually exclusive
4166        // for now — the fused-pair path doesn't carry per-token φ.
4167        let mut router = if mtp.is_none() {
4168            self.dyn_router.take()
4169        } else {
4170            None
4171        };
4172        let mut reuse_from = reuse_from;
4173        if let Some(r) = &mut router {
4174            r.reset(); // active=backbone, matching a fresh overlay
4175            self.dyn_phi_seen = 0; // fresh φ EMA per generation
4176            if self.dyn_active.is_some() {
4177                // A real switch back to the backbone invalidates the
4178                // cache the reuse key was computed against.
4179                let _ = self.set_active_skill(None);
4180                reuse_from = 0;
4181                self.last_prefill_tokens = input_ids.len();
4182            }
4183        }
4184
4185        let mut all_ids = input_ids.to_vec();
4186        let mut generated = 0usize;
4187        let mut finish_reason = "max_tokens".to_string();
4188        let mut drafted = 0usize;
4189        let mut accepted = 0usize;
4190        // DeepSeek-V4's draft quality is strongly content-dependent.  Two
4191        // consecutive paid rounds with no extra token put it on a bounded
4192        // cooldown; predictable text keeps batching, ordinary prose falls
4193        // back to the exact walk instead of paying a slow draft forever.
4194        // Local to one generation so one difficult request cannot poison the
4195        // next one, and deliberately automatic — this is not a user knob.
4196        let mut dsv4_spec_bad = 0usize;
4197        let mut dsv4_spec_retry_at = 0usize;
4198        let mut confidence: Vec<f32> = Vec::new();
4199        let trace_on = self.trace;
4200        let calib_temp = self.calib_temp;
4201        let mut traces: Vec<TokenTrace> = Vec::new();
4202
4203        // ── Prefill: forward each prompt token once, KEEP the last hidden.
4204        //    Dense prefill runs in fused pairs (weights streamed once per
4205        //    two positions — bit-identical to sequential, proven by the
4206        //    pair tests). With MTP: warm the draft head on
4207        //    (hidden_p, token_{p+1}) pairs.
4208        let mut hidden = vec![0.0f32; self.hidden_size];
4209        let mut pos = reuse_from;
4210        // lm_head-in-graph is only sound when the very next logits
4211        // consumer is this loop's own (MTP and skill routing interleave
4212        // other forwards / can swap lm_head between forward and sample).
4213        // CMF_GPU_LMHEAD=0 keeps lm_head off the graph: the token reads back
4214        // the 8 KB hidden instead of ~1 MB of logits, and the head runs on
4215        // the host. A probe for how much of the graph's fixed per-token cost
4216        // is the logits readback (the layer sweep puts that fixed part at
4217        // 3.88 ms of an 18.5 ms frame).
4218        let fuse_lm = mtp.is_none()
4219            && router.is_none()
4220            && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
4221        self.graph_logits = None;
4222        self.graph_want_logits = false;
4223        let _tpf = std::time::Instant::now();
4224        let batch_k = self.generation_batch_k();
4225        if let Some(rows) = prompt_rows {
4226            let hs = self.hidden_size;
4227            let chunk = self.prefill_chunk().max(1);
4228            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4229                let end = (pos + chunk).min(input_ids.len());
4230                let hb = match self.prefill_input_rows(
4231                    PrefillIn::Hidden(&rows[pos * hs..end * hs]), pos, task_mask,
4232                ) {
4233                    Ok(hb) => hb,
4234                    Err(err) => {
4235                        self.finish_generation(&mut mtp, &mut router, true);
4236                        return Err(err);
4237                    }
4238                };
4239                if mimo_spec { self.mimo_note_rows(&hb, pos); }
4240                hidden.copy_from_slice(&hb[hb.len() - hs..]);
4241                pos = end;
4242            }
4243        }
4244        // DeepSeek-V4 owns a separate hyper-connection stack. Route it
4245        // before the generic prefill choices: those correctly reject an
4246        // empty `weights.layers`, but their final per-position fallback used
4247        // to consume the whole prompt before `dsv4::forward_chunk` could see
4248        // it. The batch implementation therefore existed without a live
4249        // production entry point.
4250        //
4251        // Bounded chunks preserve cancellation responsiveness. Only the
4252        // prompt's final chunk asks for logits; every earlier head projection
4253        // would produce 129 280 values that no caller reads.
4254        while self.qwen4_exp.is_some()
4255            && mtp.is_none()
4256            && pos < input_ids.len()
4257            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4258        {
4259            // The device path takes several prompt tokens per layer frame;
4260            // the host path runs them one by one inside the same call.
4261            let end = (pos + crate::qwen4_exp::prefill_chunk()).min(input_ids.len());
4262            let want_logits = end == input_ids.len();
4263            let mut lg = Vec::new();
4264            if let Some(b) = &mut self.qwen4_exp {
4265                crate::qwen4_exp::forward_tokens(
4266                    &b.0,
4267                    &b.1,
4268                    &b.2,
4269                    &mut b.3,
4270                    &input_ids[pos..end],
4271                    pos,
4272                    &self.inv_freq,
4273                    self.pool.as_deref(),
4274                    &mut lg,
4275                    want_logits,
4276                );
4277            }
4278            if want_logits {
4279                self.graph_logits = Some(lg);
4280            }
4281            pos = end;
4282            hidden.fill(0.0);
4283        }
4284        while self.dsv4.is_some()
4285            && mtp.is_none()
4286            && pos < input_ids.len()
4287            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4288        {
4289            let end = (pos + prefill_chunk()).min(input_ids.len());
4290            let ids: Vec<u32> = input_ids[pos..end].to_vec();
4291            let mut lg = Vec::new();
4292            if let Some(b) = &mut self.dsv4 {
4293                let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
4294                crate::dsv4::forward_chunk(
4295                    g,
4296                    layers,
4297                    &cfg,
4298                    st,
4299                    &ids,
4300                    pos,
4301                    &self.inv_freq,
4302                    self.pool.as_deref(),
4303                    &mut lg,
4304                    end == input_ids.len(),
4305                );
4306            }
4307            if end == input_ids.len() {
4308                self.graph_logits = Some(lg);
4309            }
4310            pos = end;
4311            hidden = vec![0.0; self.hidden_size];
4312        }
4313        let dsv41_prefill = self.dsv41_prefill.take();
4314        while self.dsv41.is_some()
4315            && mtp.is_none()
4316            && pos < input_ids.len()
4317            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4318        {
4319            let end = (pos + prefill_chunk()).min(input_ids.len());
4320            let ids: Vec<u32> = input_ids[pos..end].to_vec();
4321            let mut lg = Vec::new();
4322            if let Some(b) = &mut self.dsv41 {
4323                let (g, layers, cfg, st) = (&b.0, &b.1, &b.2, &mut b.3);
4324                if let Some((embeddings, participates)) = dsv41_prefill.as_ref() {
4325                    crate::dsv41::forward_chunk_masked_with_embeddings(
4326                        g,
4327                        layers,
4328                        cfg,
4329                        st,
4330                        &ids,
4331                        pos,
4332                        &embeddings[pos..end],
4333                        &participates[pos..end],
4334                        self.pool.as_deref(),
4335                        &mut lg,
4336                    );
4337                } else {
4338                    crate::dsv41::forward_chunk(
4339                        g,
4340                        layers,
4341                        cfg,
4342                        st,
4343                        &ids,
4344                        pos,
4345                        self.pool.as_deref(),
4346                        &mut lg,
4347                    );
4348                }
4349            }
4350            if end == input_ids.len() {
4351                self.graph_logits = Some(lg);
4352            }
4353            pos = end;
4354            hidden = vec![0.0; self.hidden_size];
4355        }
4356        // With dynamic routing, prefill sequentially so the φ hook fires
4357        // over the PROMPT — the router enters decode with a warm φ (the
4358        // fused-pair path skips the per-layer φ capture). o1 layers
4359        // collect their query trace in both the single and pair paths.
4360        let dyn_prefill = router.is_some();
4361        // Optional bounded calibration prefix for generation.  The normal
4362        // O(1) path seals after the full prompt; this explicit knob instead
4363        // runs only the requested prefix through exact attention, seals the
4364        // Nyström state, and streams the rest of the prompt through the same
4365        // O(1) step used by decode.  It keeps the O(1) layers' Q trace and
4366        // temporary full KV bounded by the prefix while leaving the default
4367        // full-prompt quality profile untouched.
4368        let o1_prefill_limit = o1_prefill
4369            .and_then(|requested| self.o1_effective_boundary(requested))
4370            .map(|boundary| boundary.min(input_ids.len()));
4371        let mut o1_sealed = false;
4372        if let Some(limit) = o1_prefill_limit {
4373            // Reuse the exact batched prefix machinery when available; it
4374            // records the same per-position Q trace as the full prefill.
4375            if self.can_prefill_batched() && limit > 2 {
4376                let chunk = self.prefill_chunk();
4377                let hs = self.hidden_size;
4378                while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4379                    let end = (pos + chunk).min(limit);
4380                    let hb = self.prefill_batch(&input_ids[pos..end], pos);
4381                    hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4382                    pos = end;
4383                }
4384            } else {
4385                while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4386                    hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, None);
4387                    pos += 1;
4388                }
4389            }
4390            if pos >= limit {
4391                o1_sealed = match self.o1_seal_checked() {
4392                    Ok(sealed) => sealed,
4393                    Err(err) => {
4394                        self.finish_generation(&mut mtp, &mut router, true);
4395                        return Err(err);
4396                    }
4397                };
4398                tracing::info!(
4399                    "o1 bounded prompt prefix: requested={} effective={} processed={} of {} token(s)",
4400                    o1_prefill.unwrap_or(0),
4401                    self.o1_effective_boundary(o1_prefill.unwrap_or(0))
4402                        .unwrap_or(limit),
4403                    limit,
4404                    input_ids.len()
4405                );
4406            }
4407        }
4408        // q1 hybrids on Metal: the per-position GPU token graph beats
4409        // the CPU chunk-GEMM (whose wall is the sequential scalar GDN
4410        // recurrence), so prefill goes position-by-position through the
4411        // same graph as decode. Pure-attention models keep the batched
4412        // path — there the chunk-GEMM amortization wins.
4413        let graph_prefill = self.graph_prefill_preferred();
4414        // Native Metal, q4tp GDN hybrids: the prompt through the b-row
4415        // rows graph — projections as GEMMs over up to 512 positions, the
4416        // GDN recurrence in registers on the device, K/V rows appended by
4417        // the chunk — instead of one token-graph submit per position (the
4418        // 27B: 8 tok/s → GEMM-bound). The MTP warm-up rows come out of one
4419        // batched run of the block per chunk. Any refusal leaves the rest
4420        // of the prompt to the sequential paths below.
4421        #[cfg(target_os = "macos")]
4422        if task_mask.is_none()
4423            && !dyn_prefill
4424            && (crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in())
4425            && crate::gpu::enabled_here()
4426            && self.gdn_cfg.is_some()
4427            && self.g3n.is_none()
4428            && input_ids.len() > 8
4429            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
4430            && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
4431        {
4432            let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
4433                .ok()
4434                .and_then(|v| v.parse().ok())
4435                .filter(|&v| (16..=512).contains(&v))
4436                .unwrap_or(256);
4437            let hs = self.hidden_size;
4438            let _tp = std::time::Instant::now();
4439            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4440                let end = (pos + chunk).min(input_ids.len());
4441                let hb = match self.prefill_batch_metal(&input_ids[pos..end], pos) {
4442                    MetalPrefillOutcome::Completed(hb) => hb,
4443                    MetalPrefillOutcome::Declined => break,
4444                    MetalPrefillOutcome::Failed => {
4445                        self.finish_generation(&mut mtp, &mut router, true);
4446                        return Err("ordinary Metal prefill failed after admission".into());
4447                    }
4448                };
4449                if let Some(m) = &mut mtp {
4450                    let n_pairs = if end < input_ids.len() {
4451                        end - pos
4452                    } else {
4453                        end - pos - 1
4454                    };
4455                    if n_pairs > 0 {
4456                        let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
4457                            .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
4458                            .collect();
4459                        if !self.mtp_warm_batch_metal(m, &pairs, pos) {
4460                            for (j, (h, t)) in pairs.iter().enumerate() {
4461                                let h = h.to_vec();
4462                                let _ = self.mtp_step(m, &h, *t, pos + j);
4463                            }
4464                        }
4465                    }
4466                }
4467                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4468                pos = end;
4469            }
4470            if std::env::var("CMF_PREFILL_PROF").is_ok() {
4471                eprintln!(
4472                    "metal-prefill: {} of {} tokens in {:.1} ms",
4473                    pos,
4474                    input_ids.len(),
4475                    _tp.elapsed().as_secs_f64() * 1e3
4476                );
4477            }
4478        }
4479        self.mimo_moe_prepare();
4480        // A MoE stack larger than the card (MiMo-V2 q4tp on 96 GB): the
4481        // batched wgpu graph runs the device prefix of every chunk — its
4482        // experts resident — and the host's batched layer walk finishes
4483        // the chunk. Any refusal leaves the rest of the prompt to the
4484        // chunked prefill below.
4485        #[cfg(not(target_os = "macos"))]
4486        if task_mask.is_none()
4487            && !dyn_prefill
4488            && !graph_prefill
4489            && mtp.is_none()
4490            && o1_prefill.is_none()
4491            && !self.o1_active()
4492            && input_ids.len() > 2
4493            && self.batch_prefix_prefill()
4494        {
4495            let chunk = self.prefill_chunk().max(1);
4496            let hs = self.hidden_size;
4497            let t_bp = std::time::Instant::now();
4498            let pos0 = pos;
4499            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4500                let end = (pos + chunk).min(input_ids.len());
4501                let bk = end - pos;
4502                let mut hiddens = vec![0f32; bk * hs];
4503                for (j, &id) in input_ids[pos..end].iter().enumerate() {
4504                    hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4505                }
4506                let positions: Vec<usize> = (pos..end).collect();
4507                let mut run = 0usize;
4508                let outcome = self.try_batch_graph_wgpu_prefix(
4509                    &mut hiddens,
4510                    &positions,
4511                    bk,
4512                    None,
4513                    Some(&mut run),
4514                );
4515                match outcome {
4516                    crate::gpu::BatchGraphOutcome::Completed => {
4517                        let hb = if run < self.num_layers {
4518                            self.prefill_batch_span(
4519                                PrefillIn::Hidden(&hiddens),
4520                                pos,
4521                                None,
4522                                run,
4523                                self.num_layers,
4524                            )
4525                        } else {
4526                            hiddens
4527                        };
4528                        if mimo_spec {
4529                            self.mimo_note_rows(&hb, pos);
4530                        }
4531                        hidden.copy_from_slice(&hb[(bk - 1) * hs..]);
4532                        pos = end;
4533                    }
4534                    crate::gpu::BatchGraphOutcome::Failed => {
4535                        self.finish_generation(&mut mtp, &mut router, true);
4536                        return Err("batched prefix prefill failed after admission".into());
4537                    }
4538                    crate::gpu::BatchGraphOutcome::Declined => {
4539                        // Earlier chunks left their prefix rows on the
4540                        // device only: the host walk below needs them.
4541                        #[cfg(feature = "gpu")]
4542                        if pos > pos0 {
4543                            self.pull_lagging_host_kv(0, self.num_layers, pos);
4544                        }
4545                        break;
4546                    }
4547                }
4548            }
4549            if std::env::var("CMF_PREFILL_PROF").is_ok() {
4550                eprintln!(
4551                    "batch-prefix prefill: {} of {} tokens in {:.1} ms",
4552                    pos - pos0,
4553                    input_ids.len(),
4554                    t_bp.elapsed().as_secs_f64() * 1e3
4555                );
4556            }
4557        }
4558        if task_mask.is_none()
4559            && !dyn_prefill
4560            && !graph_prefill
4561            && self.can_prefill_batched()
4562            && self.g3n.is_none()
4563            && o1_prefill.is_none()
4564            && input_ids.len() > 2
4565        {
4566            // Production prefill = the same chunked prefill-GEMM that
4567            // bench/PPL measure (roadmap §3 P0: generation used to warm
4568            // the prompt with the slower pair path — the published
4569            // prefill number didn't match real TTFT). MTP warm-up reads
4570            // each position's hidden straight from the chunk result.
4571            let chunk = self.prefill_chunk();
4572            let hs = self.hidden_size;
4573            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4574                let end = (pos + chunk).min(input_ids.len());
4575                let hb = self.prefill_batch(&input_ids[pos..end], pos);
4576                if mimo_spec {
4577                    self.mimo_note_rows(&hb, pos);
4578                }
4579                if let Some(m) = &mut mtp {
4580                    let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4581                        .ok()
4582                        .and_then(|v| v.parse().ok())
4583                        .unwrap_or(0);
4584                    for p in pos..end {
4585                        if p + 1 < input_ids.len() {
4586                            if probe >= 1 && p + 2 < input_ids.len() {
4587                                // Teacher-forced chain acceptance (see the
4588                                // tail loop's twin): the warm-up row stays,
4589                                // the chain's rows roll back.
4590                                let (d1, mut hx) = self.mtp_step_h(
4591                                    m,
4592                                    &hb[(p - pos) * hs..(p - pos + 1) * hs],
4593                                    input_ids[p + 1],
4594                                    p,
4595                                );
4596                                let mut ok = d1 == input_ids[p + 2];
4597                                Self::chain_probe_note(0, ok);
4598                                let mut d_prev = d1;
4599                                let mut extra = 0usize;
4600                                for j in 1..probe {
4601                                    if p + 2 + j >= input_ids.len() {
4602                                        break;
4603                                    }
4604                                    let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
4605                                    extra += 1;
4606                                    ok = ok && dj == input_ids[p + 2 + j];
4607                                    Self::chain_probe_note(j, ok);
4608                                    d_prev = dj;
4609                                    hx = hj;
4610                                }
4611                                m.kv.truncate_last(extra);
4612                            } else {
4613                                let _ = self.mtp_step(
4614                                    m,
4615                                    &hb[(p - pos) * hs..(p - pos + 1) * hs],
4616                                    input_ids[p + 1],
4617                                    p,
4618                                );
4619                            }
4620                        }
4621                    }
4622                }
4623                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4624                pos = end;
4625            }
4626        }
4627        let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
4628        if task_mask.is_none()
4629            && !dyn_prefill
4630            && !graph_prefill
4631            && !pair_off
4632            && self.pair_supported()
4633            && o1_prefill.is_none()
4634        {
4635            while pos + 1 < input_ids.len()
4636                && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4637            {
4638                let e1 = self.embed_single(input_ids[pos]);
4639                let e2 = self.embed_single(input_ids[pos + 1]);
4640                let (h1, h2) = self.forward_pair(&e1, &e2, pos);
4641                if mimo_spec {
4642                    self.mimo_note_rows(&h1, pos);
4643                    self.mimo_note_rows(&h2, pos + 1);
4644                }
4645                // Both prefill tokens are real → commit lane-2 states.
4646                self.commit_linear_scratch();
4647                if let Some(m) = &mut mtp {
4648                    let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
4649                    if pos + 2 < input_ids.len() {
4650                        let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4651                            .ok()
4652                            .and_then(|v| v.parse().ok())
4653                            .unwrap_or(0);
4654                        if probe >= 1 && pos + 3 < input_ids.len() {
4655                            // Same teacher-forced chain table as the tail
4656                            // loop below, fed from the pair path that owns
4657                            // most prefill positions.
4658                            let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
4659                            let mut ok = d1 == input_ids[pos + 3];
4660                            Self::chain_probe_note(0, ok);
4661                            let mut d_prev = d1;
4662                            let mut extra = 0usize;
4663                            for j in 1..probe {
4664                                if pos + 3 + j >= input_ids.len() {
4665                                    break;
4666                                }
4667                                let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
4668                                extra += 1;
4669                                ok = ok && dj == input_ids[pos + 3 + j];
4670                                Self::chain_probe_note(j, ok);
4671                                d_prev = dj;
4672                                hx = hj;
4673                            }
4674                            m.kv.truncate_last(extra);
4675                        } else {
4676                            let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
4677                        }
4678                    }
4679                }
4680                hidden = h2;
4681                pos += 2;
4682            }
4683        }
4684        // Batched GPU prefill for the wgpu decode graph (GDN hybrids): K prompt
4685        // positions per submit — projections/FFN as GEMMs (weight once per K),
4686        // attention/GDN looped inside — instead of one whole-graph submit per
4687        // position. Falls through to the per-position graph on any refusal.
4688        // Batched prefill is opt-in (CMF_BATCH_K>0). Default 0 = per-position
4689        // graph prefill. (Steady-state decode is provably identical either way —
4690        // token-graph submit and lm_head both unchanged — so this only trades
4691        // prefill wall.)
4692        // A bounded O(1) prefix is the one post-seal prompt interval: only
4693        // admit its batch when the device O(1) route is explicitly enabled and
4694        // every sealed layer exposes a portable view. The same batch size and
4695        // refusal behavior remain the ordinary controls/comparator.
4696        let o1_batch_ready = o1_sealed
4697            && o1_prefill.is_some()
4698            && mtp.is_none()
4699            && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
4700            && (0..self.num_layers).all(|li| {
4701                let cache = &self.kv_cache.layers[self.phys_layer(li)];
4702                cache.o1.is_none() || cache.o1_views().is_some()
4703            });
4704        // The ordinary graph-prefill route can share each completed trunk
4705        // chunk with an attached MTP head.  Keep chain probing on its
4706        // established per-position path: the probe deliberately needs every
4707        // teacher-forced draft row and its rollback table.
4708        let mtp_batch_prefill = mtp.is_some()
4709            && graph_prefill
4710            && task_mask.is_none()
4711            && !dyn_prefill
4712            && !self.o1_active()
4713            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err();
4714        if batch_k > 0
4715            && (graph_prefill || o1_batch_ready)
4716            && task_mask.is_none()
4717            && (!self.o1_active() || o1_batch_ready)
4718            && (mtp.is_none() || mtp_batch_prefill)
4719            && !dyn_prefill
4720            && pos + 1 < input_ids.len()
4721        {
4722            let hs = self.hidden_size;
4723            let chunk = batch_k;
4724            while pos < input_ids.len() {
4725                let end = (pos + chunk).min(input_ids.len());
4726                let bk = end - pos;
4727                let mut hiddens = vec![0f32; bk * hs];
4728                for (j, &id) in input_ids[pos..end].iter().enumerate() {
4729                    hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4730                }
4731                let positions: Vec<usize> = (pos..end).collect();
4732                let t_chunk = std::time::Instant::now();
4733                let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
4734                let ok_b = outcome == crate::gpu::BatchGraphOutcome::Completed;
4735                if std::env::var("CMF_GRAPH_PROF").is_ok() {
4736                    let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
4737                    eprintln!(
4738                        "batch-chunk: phase=prompt mode={} k={bk} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
4739                        if o1_batch_ready {
4740                            "o1"
4741                        } else if mtp_batch_prefill {
4742                            "ordinary_mtp"
4743                        } else {
4744                            "ordinary"
4745                        },
4746                        bk as f64 / (ms / 1000.0)
4747                    );
4748                }
4749                {
4750                    use std::sync::atomic::{AtomicBool, Ordering};
4751                    static SAID: AtomicBool = AtomicBool::new(false);
4752                    if !SAID.swap(true, Ordering::Relaxed) {
4753                        if ok_b {
4754                            tracing::info!(
4755                                "batched prefill: ACTIVE mode={} (k={bk})",
4756                                if o1_batch_ready {
4757                                    "o1"
4758                                } else if mtp_batch_prefill {
4759                                    "ordinary_mtp"
4760                                } else {
4761                                    "ordinary"
4762                                }
4763                            );
4764                        } else {
4765                            tracing::warn!("batched prefill {:?} — per-position graph", outcome);
4766                        }
4767                    }
4768                }
4769                if ok_b {
4770                    if mimo_spec {
4771                        self.mimo_note_rows(&hiddens, pos);
4772                    }
4773                    if mtp_batch_prefill {
4774                        let n_pairs = mtp_prefill_pair_count(pos, end, input_ids.len());
4775                        if n_pairs > 0 {
4776                            // `hiddens` is owned by this chunk, so materialize
4777                            // row slices before borrowing the detached MTP
4778                            // module.  The last prompt row has no successor;
4779                            // the helper above is the single source of that
4780                            // boundary rule.
4781                            let rows: Vec<Vec<f32>> = (0..n_pairs)
4782                                .map(|j| hiddens[j * hs..(j + 1) * hs].to_vec())
4783                                .collect();
4784                            let pairs: Vec<(&[f32], u32)> = rows
4785                                .iter()
4786                                .enumerate()
4787                                .map(|(j, row)| (row.as_slice(), input_ids[pos + j + 1]))
4788                                .collect();
4789                            if std::env::var("CMF_GRAPH_PROF").is_ok() {
4790                                eprintln!(
4791                                    "mtp-warm: phase=prompt mode=ordinary_mtp first_pos={} pairs={} last_pos={}",
4792                                    pos,
4793                                    n_pairs,
4794                                    pos + n_pairs - 1,
4795                                );
4796                            }
4797                            let warm_error = if let Some(m) = mtp.as_mut() {
4798                                self.mtp_warm_prefill_pairs(m, &pairs, pos).err()
4799                            } else {
4800                                None
4801                            };
4802                            if let Some(err) = warm_error {
4803                                // The trunk batch was already admitted.  A
4804                                // failed MTP warm-up therefore clears both
4805                                // mirrors and exits; continuing would pair a
4806                                // current trunk state with a stale MTP cache.
4807                                self.finish_generation(&mut mtp, &mut router, true);
4808                                return Err(err.to_string());
4809                            }
4810                        }
4811                    }
4812                    hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
4813                    pos = end;
4814                } else if outcome == crate::gpu::BatchGraphOutcome::Failed {
4815                    // A failed batch may have advanced a device recurrent
4816                    // state (ordinary GDN or sealed O(1)). A CPU fallback
4817                    // would then observe stale accumulators, so clear the
4818                    // request state and make the failure explicit.
4819                    self.finish_generation(&mut mtp, &mut router, true);
4820                    return Err(if o1_batch_ready {
4821                        "sealed O(1) batch graph failed after admission".to_string()
4822                    } else {
4823                        "ordinary recurrent batch graph failed after admission".to_string()
4824                    });
4825                } else {
4826                    break; // unsupported → per-position graph handles the rest
4827                }
4828            }
4829        }
4830        // Resident Embryo graph: the prompt in chunks of one submit each
4831        // instead of one whole-graph submit per position; the last chunk
4832        // carries the logits exactly as the per-position walk would.
4833        if graph_prefill
4834            && task_mask.is_none()
4835            && mtp.is_none()
4836            && !dyn_prefill
4837            && pos == 0
4838            && input_ids.len() > 1
4839            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4840        {
4841            if let Some(lg) = self.embryo_prefill_chunked(input_ids, 0) {
4842                self.graph_logits = Some(lg);
4843                hidden = vec![0.0; self.hidden_size];
4844                pos = input_ids.len();
4845            }
4846        }
4847        while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4848            self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
4849            hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
4850            if mimo_spec {
4851                self.mimo_note_rows(&hidden, pos);
4852            }
4853            if let Some(m) = &mut mtp {
4854                if pos + 1 < input_ids.len() {
4855                    // `CMF_MTP_CHAIN_PROBE=k`: teacher-forced acceptance of a
4856                    // CHAINED draft — iterate the head on its own hidden k
4857                    // deep and score every depth against the prompt's real
4858                    // continuation. The economics of a k-token speculative
4859                    // round stand or fall on this table.
4860                    let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4861                        .ok()
4862                        .and_then(|v| v.parse().ok())
4863                        .unwrap_or(0);
4864                    if probe >= 1 && pos + 2 < input_ids.len() {
4865                        let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
4866                        let mut ok = d1 == input_ids[pos + 2];
4867                        Self::chain_probe_note(0, ok);
4868                        let mut d_prev = d1;
4869                        let mut extra = 0usize;
4870                        for j in 1..probe {
4871                            if pos + 2 + j >= input_ids.len() {
4872                                break;
4873                            }
4874                            let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
4875                            extra += 1;
4876                            ok = ok && dj == input_ids[pos + 2 + j];
4877                            Self::chain_probe_note(j, ok);
4878                            d_prev = dj;
4879                            hx = hj;
4880                        }
4881                        // The chain's rows are speculation, not the prompt —
4882                        // keep only the warmup row the plain path would add.
4883                        m.kv.truncate_last(extra);
4884                    } else {
4885                        let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
4886                    }
4887                }
4888            }
4889            pos += 1;
4890        }
4891        if std::env::var("CMF_PREFILL_PROF").is_ok() {
4892            eprintln!(
4893                "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
4894                input_ids.len(),
4895                _tpf.elapsed().as_secs_f64() * 1000.0
4896            );
4897        }
4898        if self
4899            .graph_failed
4900            .swap(false, std::sync::atomic::Ordering::Relaxed)
4901        {
4902            // MTP is detached for speculative generation.  Restore the
4903            // module before returning the terminal graph error; otherwise a
4904            // failed request would silently remove the head from a pooled
4905            // pipeline and the next request would lose its configured route.
4906            self.finish_generation(&mut mtp, &mut router, true);
4907            return Err("GPU token graph failed during prefill".to_string());
4908        }
4909        // Cancelled mid-prefill: the cache holds a partial prompt —
4910        // drop the reuse history and return an empty generation.
4911        if self
4912            .cancel
4913            .swap(false, std::sync::atomic::Ordering::Relaxed)
4914        {
4915            // A cancelled prefill can already have advanced the device
4916            // mirror. Drop the whole partial sequence so a pooled pipeline
4917            // cannot carry that state into its next request.
4918            self.finish_generation(&mut mtp, &mut router, true);
4919            return Ok(GenerateResult {
4920                text: String::new(),
4921                token_ids: Vec::new(),
4922                prompt_tokens: input_ids.len(),
4923                tokens_generated: 0,
4924                finish_reason: "cancelled".to_string(),
4925                mtp_drafted: 0,
4926                mtp_accepted: 0,
4927                token_confidence: Vec::new(),
4928                traces: Vec::new(),
4929            });
4930        }
4931
4932        // Prompt absorbed → freeze the o1 layers' skeletons; from here
4933        // every decode step on those layers is O(W + m·dv + m²).
4934        if !o1_sealed {
4935            match self.o1_seal_checked() {
4936                Ok(_) => {}
4937                Err(err) => {
4938                    self.finish_generation(&mut mtp, &mut router, true);
4939                    return Err(err);
4940                }
4941            }
4942        }
4943
4944        // Commit one token: push, check EOS, stream. Returns false = stop.
4945        macro_rules! commit {
4946            ($id:expr) => {{
4947                all_ids.push($id);
4948                generated += 1;
4949                self.note_draft_id($id);
4950                if self.tokenizer.is_eos($id) && !self.ignore_eos {
4951                    finish_reason = "stop".to_string();
4952                    false
4953                } else {
4954                    let token_text = self.tokenizer.decode_token($id);
4955                    let mut go = true;
4956                    if let Some(ref mut cb) = on_token {
4957                        if !cb(&token_text) {
4958                            finish_reason = "cancelled".to_string();
4959                            go = false;
4960                        }
4961                    }
4962                    go
4963                }
4964            }};
4965        }
4966
4967        // Speculation is decided by MEASUREMENT, not by an acceptance
4968        // model. A k=4 round costs ~3.8 plain tokens on the 5090 (draft
4969        // 6.6 + verify 66.6 + commit 4.8 ms against a 20.6 ms token), so it
4970        // pays only when the head lands ~2.8 of 4 — predictable text (code,
4971        // structured output) does, free prose often does not, and the
4972        // ratio at which the two cross depends on the card and the context
4973        // depth. So: four speculative rounds timed, then eight plain
4974        // tokens timed, and the faster arm runs until a re-check 256
4975        // tokens later (context growth moves the balance). The trial
4976        // costs at most a few tokens of the slower arm per 256.
4977        let mut spec_trial = SpecTrial::Spec {
4978            t0: std::time::Instant::now(),
4979            gen0: generated,
4980            rounds: 0,
4981        };
4982        // The token-count proxy prices a round at ~1.9 plain tokens. That
4983        // holds for the Metal rounds whose cost was measured — greedy and
4984        // the sparse sampling chain — so an expensive round (the dense
4985        // chain, reachable only by `CMF_GRAPH_SPEC_SAMPLE=1`) still times
4986        // the plain path before it decides.
4987        let mut spec_mon = SpecMon {
4988            metal: graph_spec && crate::gpu::q1_force() && spec_cheap_round,
4989            ..SpecMon::default()
4990        };
4991        let mut spec_watchdog_off = false;
4992        // CMF_GRAPH_SPEC_TIME: the round walls so far (round 1 excluded —
4993        // it pays the scratch), for the outlier test on each new one
4994        let mut spec_walls: Vec<f32> = Vec::new();
4995        // ... and the end of the last round: the host time between rounds
4996        // (token commits, streaming, the loop top) is printed at level 2
4997        let mut spec_round_end: Option<std::time::Instant> = None;
4998        if mimo_spec {
4999            if let Ok(path) = std::env::var("CMF_MIMO_MTP_PROBE") {
5000                if let Some(mut st) = self.mimo_mtp.take() {
5001                    self.mimo_mtp_probe(&mut st, input_ids, &path);
5002                    self.mimo_mtp = Some(st);
5003                }
5004            }
5005        }
5006        // ── Decode ──
5007        let mut next_pos = input_ids.len();
5008        'decode: while generated < max_tokens {
5009            if self
5010                .graph_failed
5011                .swap(false, std::sync::atomic::Ordering::Relaxed)
5012            {
5013                // Keep the detached MTP module attached after a terminal
5014                // graph error so the pipeline can be reused for a fresh
5015                // sequence.  `clear_sequence_state` only clears mirrors and
5016                // host KV; it cannot recover a module dropped here.
5017                self.finish_generation(&mut mtp, &mut router, true);
5018                return Err("GPU token graph failed during decode".to_string());
5019            }
5020            if self
5021                .cancel
5022                .swap(false, std::sync::atomic::Ordering::Relaxed)
5023            {
5024                finish_reason = "cancelled".to_string();
5025                break 'decode;
5026            }
5027            // A rejected speculative draft already drew this position's
5028            // token from the residual distribution (graph_spec_step); it
5029            // is committed as-is — sampling again from the row's logits
5030            // would bias the stream toward the target's mode.
5031            if mimo_spec && next_pos > 0 {
5032                // Every path leaves `hidden` = the backbone output at
5033                // next_pos-1; the draft layers read it (idempotent).
5034                self.mimo_note_rows(&hidden, next_pos - 1);
5035            }
5036            let forced = self.spec_forced.take();
5037            let mut logits = match (forced, self.graph_logits.take()) {
5038                (Some(_), _) => Vec::new(),
5039                (None, Some(lg)) => lg,
5040                (None, None) => {
5041                    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Head);
5042                    inference::rms_norm_into(
5043                        &hidden,
5044                        &self.weights.final_norm,
5045                        self.rms_eps,
5046                        self.norm_style,
5047                        &mut self.ws.n1,
5048                    );
5049                    self.lm_head_forward(&self.ws.n1)
5050                }
5051            };
5052            // CMF_LOGIT_DUMP=<path>: the first decode step's hidden + logits
5053            // as raw f32 (hidden first) — cross-backend numerics diffing.
5054            if generated
5055                == std::env::var("CMF_LOGIT_DUMP_STEP")
5056                    .ok()
5057                    .and_then(|v| v.parse().ok())
5058                    .unwrap_or(0)
5059            {
5060                if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
5061                    let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
5062                    for v in hidden.iter().chain(logits.iter()) {
5063                        bytes.extend_from_slice(&v.to_le_bytes());
5064                    }
5065                    if let Err(e) = std::fs::write(&path, &bytes) {
5066                        eprintln!("logit dump: failed to write {path}: {e}");
5067                        self.finish_generation(&mut mtp, &mut router, true);
5068                        return Err(format!("logit dump write failed: {e}"));
5069                    }
5070                }
5071            }
5072            // CMF_LOGIT_DUMP_ALL=<dir>: every decode step's logits as raw
5073            // f32, `<dir>/step{n:05}.f32` — step-by-step backend diffing
5074            // (a greedy run on two backends compares until they diverge).
5075            if let Ok(dir) = std::env::var("CMF_LOGIT_DUMP_ALL") {
5076                if !logits.is_empty() {
5077                    let path = std::path::Path::new(&dir).join(format!("step{generated:05}.f32"));
5078                    let bytes: Vec<u8> = logits.iter().flat_map(|v| v.to_le_bytes()).collect();
5079                    if let Err(e) =
5080                        std::fs::create_dir_all(&dir).and_then(|_| std::fs::write(&path, &bytes))
5081                    {
5082                        eprintln!("logit dump: failed to write {}: {e}", path.display());
5083                    }
5084                }
5085            }
5086            let t_next = match forced {
5087                Some(c) => c,
5088                None => {
5089                    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Sampler);
5090                    sampler::sample_with_scratch_pool(
5091                        &logits,
5092                        &self.sampler_config,
5093                        self.sampler_config.penalty_past(&all_ids, bounded_native),
5094                        &mut self.rng,
5095                        &mut self.sampler_scratch,
5096                        self.pool.as_deref(),
5097                    )
5098                }
5099            };
5100            if self.confidence_on {
5101                confidence.push(if logits.is_empty() {
5102                    0.0
5103                } else {
5104                    sampler::top1_prob_pool(
5105                        self.pool.as_deref(),
5106                        &mut self.sampler_scratch,
5107                        &logits,
5108                        t_next,
5109                        calib_temp,
5110                    )
5111                });
5112            }
5113            if !logits.is_empty() {
5114                attention::recycle_buf(&mut logits);
5115            }
5116            if trace_on {
5117                // active_skill = the overlay in force while this token was
5118                // generated; recon/switched are filled after the post-emit
5119                // routing eval below (freshest coherence for this token).
5120                let skill = router.as_ref().and_then(|r| r.active_id());
5121                traces.push(TokenTrace {
5122                    t: generated,
5123                    token_id: t_next,
5124                    confidence: confidence.last().copied().unwrap_or(0.0),
5125                    active_skill: skill,
5126                    recon: None,
5127                    switched: false,
5128                });
5129            }
5130            if !commit!(t_next) {
5131                break 'decode;
5132            }
5133            if generated >= max_tokens {
5134                break 'decode;
5135            }
5136
5137            if self.dsv41.is_none() && self.kv_cache.needs_eviction() {
5138                // Say it ONCE, loudly: past this point the model keeps
5139                // talking but has lost half its context, and on a GDN
5140                // hybrid the graph's device state goes stale on top. The
5141                // Qwen3.8 bring-up spent a day reading this cliff as
5142                // three different model bugs.
5143                static SAID: std::sync::Once = std::sync::Once::new();
5144                SAID.call_once(|| {
5145                    tracing::warn!(
5146                        "KV cache full at {} positions — evicting half; quality \
5147                         will degrade. Raise CMF_MAX_SEQ.",
5148                        self.kv_cache.max_seq_len,
5149                    );
5150                });
5151                let keep = (self.kv_cache.max_seq_len / 2).max(1);
5152                self.kv_cache.evict(keep);
5153            }
5154
5155            // Advance the speculation trial: plain-phase accounting and
5156            // the periodic re-check happen here, on every token.
5157            if graph_spec {
5158                match spec_trial {
5159                    SpecTrial::Plain { t0, gen0 } if spec_mon.plain_done(t0, gen0, generated) => {
5160                        spec_mon.plain_ms =
5161                            t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
5162                        let keep = spec_mon.pays();
5163                        tracing::info!(
5164                            "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5165                            spec_mon.tokens,
5166                            spec_mon.round_ms,
5167                            spec_mon.plain_ms,
5168                            if keep { "speculating" } else { "plain" }
5169                        );
5170                        spec_mon.fails = 0;
5171                        spec_trial = SpecTrial::Decided {
5172                            spec: keep,
5173                            recheck_at: if keep { usize::MAX } else { generated + 128 },
5174                        };
5175                    }
5176                    SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
5177                        spec_mon.n = 0;
5178                        spec_trial = SpecTrial::Spec {
5179                            t0: std::time::Instant::now(),
5180                            gen0: generated,
5181                            rounds: 0,
5182                        };
5183                    }
5184                    _ => {}
5185                }
5186                spec_watchdog_off = matches!(
5187                    spec_trial,
5188                    SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
5189                );
5190            }
5191            // ── MiMo-V2 draft stack: draft K, verify K+1 rows in one batch ──
5192            if mimo_spec && generated + 1 < max_tokens && next_pos > 0 {
5193                let budget = max_tokens - generated - 1;
5194                if let Some(mut st) = self.mimo_mtp.take() {
5195                    let k = st.depth.min(budget);
5196                    let r = self.mimo_spec_round(&mut st, next_pos, &all_ids, k);
5197                    self.mimo_mtp = Some(st);
5198                    let r = match r {
5199                        Ok(r) => r,
5200                        Err(err) => {
5201                            self.finish_generation(&mut mtp, &mut router, true);
5202                            return Err(err);
5203                        }
5204                    };
5205                    if let Some(r) = r {
5206                        drafted += r.drafted;
5207                        accepted += r.accepted.len();
5208                        let mut stopped = false;
5209                        for &id in &r.accepted {
5210                            if self.confidence_on {
5211                                confidence.push(0.0);
5212                            }
5213                            if !commit!(id) {
5214                                stopped = true;
5215                                break;
5216                            }
5217                        }
5218                        if stopped {
5219                            break 'decode;
5220                        }
5221                        next_pos += r.accepted.len() + 1;
5222                        hidden = r.hidden;
5223                        // The loop top chooses the round's own token from
5224                        // these logits — the same sampler, same history.
5225                        self.graph_logits = Some(r.logits);
5226                        continue 'decode;
5227                    }
5228                }
5229            }
5230            // ── Qwen3.8-Flash-Next draft head: k greedy drafts from the MTP
5231            //    sidecar, one batched verify window on the device ──
5232            #[cfg(feature = "gpu")]
5233            if self.speculative
5234                && self.qwen4_exp.is_some()
5235                && task_mask.is_none()
5236                && self.sampler_config.temperature < 1e-6
5237                && generated + 1 < max_tokens
5238                && next_pos > 0
5239                && std::env::var("CMF_QWEN_MTP").as_deref() != Ok("0")
5240            {
5241                let r = match &mut self.qwen4_exp {
5242                    Some(b) => crate::qwen4_exp::spec_round(
5243                        &b.0,
5244                        &b.1,
5245                        &b.2,
5246                        &mut b.3,
5247                        next_pos,
5248                        &all_ids,
5249                        &self.inv_freq,
5250                        self.pool.as_deref(),
5251                    ),
5252                    None => None,
5253                };
5254                if let Some(r) = r {
5255                    drafted += r.drafted;
5256                    accepted += r.accepted.len();
5257                    let mut stopped = false;
5258                    for &id in &r.accepted {
5259                        if self.confidence_on {
5260                            confidence.push(0.0);
5261                        }
5262                        if !commit!(id) {
5263                            stopped = true;
5264                            break;
5265                        }
5266                    }
5267                    if stopped {
5268                        break 'decode;
5269                    }
5270                    next_pos += r.accepted.len() + 1;
5271                    hidden.fill(0.0);
5272                    self.graph_logits = Some(r.logits);
5273                    continue 'decode;
5274                }
5275            }
5276            match &mut mtp {
5277                // ── Graph speculation: chain-draft, batch-verify on device ──
5278                #[cfg(feature = "gpu")]
5279                Some(m)
5280                    if graph_spec
5281                        && !spec_watchdog_off
5282                        && generated + 1 < max_tokens
5283                        && next_pos > 0 =>
5284                {
5285                    let t_round = std::time::Instant::now();
5286                    if spec_time_level() >= 2 {
5287                        if let Some(t) = spec_round_end.take() {
5288                            eprintln!(
5289                                "spec-gap {:.2} ms (host between rounds)",
5290                                t.elapsed().as_secs_f64() * 1e3
5291                            );
5292                        }
5293                    }
5294                    spec_stamps_begin();
5295                    // device buffers allocated during this round: a
5296                    // first-touch Shared allocation is zero-filled inside
5297                    // the command buffer that uses it, which is what the
5298                    // long outlier rounds were
5299                    #[cfg(target_os = "macos")]
5300                    let allocs0 = crate::gpu_metal::IO_BUF_ALLOCS
5301                        .load(std::sync::atomic::Ordering::Relaxed);
5302                    #[cfg(not(target_os = "macos"))]
5303                    let allocs0 = 0u64;
5304                    if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
5305                        m,
5306                        &hidden,
5307                        t_next,
5308                        next_pos,
5309                        &mut drafted,
5310                        &mut accepted,
5311                        &mut all_ids,
5312                        max_tokens - generated,
5313                    ) {
5314                        next_pos = n_pos;
5315                        hidden = new_h;
5316                        let level = spec_time_level();
5317                        if level > 0 {
5318                            let wall = t_round.elapsed().as_secs_f32() * 1e3;
5319                            let stamps = spec_stamps_take();
5320                            // the running median of the rounds before this
5321                            // one (round 1 pays the scratch: not a sample)
5322                            let median = if spec_walls.len() >= 3 {
5323                                let mut s = spec_walls.clone();
5324                                s.sort_by(|a, b| a.partial_cmp(b).unwrap());
5325                                Some(s[s.len() / 2])
5326                            } else {
5327                                None
5328                            };
5329                            let outlier = median.is_some_and(|m| wall > 1.4 * m);
5330                            #[cfg(target_os = "macos")]
5331                            let allocs = crate::gpu_metal::IO_BUF_ALLOCS
5332                                .load(std::sync::atomic::Ordering::Relaxed)
5333                                - allocs0;
5334                            #[cfg(not(target_os = "macos"))]
5335                            let allocs = allocs0;
5336                            eprintln!(
5337                                "spec-round wall {wall:.1} ms → {} tokens{}{}",
5338                                extra.len() + 1,
5339                                if allocs > 0 {
5340                                    format!(" [{allocs} new device buffers]")
5341                                } else {
5342                                    String::new()
5343                                },
5344                                match (outlier, median) {
5345                                    (true, Some(m)) => format!(" OUTLIER (median {m:.1})"),
5346                                    _ => String::new(),
5347                                }
5348                            );
5349                            if level >= 2 || outlier {
5350                                let sum: f32 = stamps.iter().map(|s| s.1).sum();
5351                                eprintln!(
5352                                    "spec-stamps: {}| untracked {:.1}",
5353                                    spec_stamps_format(&stamps),
5354                                    wall - sum
5355                                );
5356                            }
5357                            if spec_mon.n >= 1 {
5358                                spec_walls.push(wall);
5359                            }
5360                        }
5361                        // One speculative round done: the monitor counts it
5362                        // (round 1 untimed — it pays the batch scratch and
5363                        // the draft mirror), and the trial advances.
5364                        spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
5365                        // the round's tokens land in `generated` below; the
5366                        // plain phase must start counting AFTER them
5367                        spec_trial = Self::spec_trial_round(
5368                            spec_trial,
5369                            &mut spec_mon,
5370                            generated + extra.len() + 1,
5371                        );
5372                        let mut stopped = false;
5373                        for &id in &extra {
5374                            if self.confidence_on {
5375                                confidence.push(0.0);
5376                            }
5377                            if !commit!(id) {
5378                                stopped = true;
5379                                break;
5380                            }
5381                        }
5382                        if stopped {
5383                            break 'decode;
5384                        }
5385                        if spec_time_level() >= 2 {
5386                            spec_round_end = Some(std::time::Instant::now());
5387                        }
5388                        continue 'decode;
5389                    }
5390                    if self
5391                        .graph_failed
5392                        .swap(false, std::sync::atomic::Ordering::Relaxed)
5393                    {
5394                        // `graph_spec_step` may have detached MTP while a
5395                        // warm-up was in flight.  Do not reinterpret its
5396                        // terminal device failure as a plain decode step;
5397                        // restore the head, clear both mirrors, and surface
5398                        // one explicit error to the caller.
5399                        self.finish_generation(&mut mtp, &mut router, true);
5400                        return Err("GPU MTP graph failed during speculative decode".to_string());
5401                    }
5402                    // Declined (batch graph refused): plain forward below —
5403                    // and a round that produced one token for the trial's
5404                    // ledger, so a graph that keeps refusing is measured out
5405                    // like a head that keeps missing (it was spinning
5406                    // forever on a file whose batch graph declines).
5407                    // A declined round is not a cheap one-token round — it
5408                    // is a verify that does not exist for this file (a
5409                    // healed q8_2f tail measured 760 drafts, 0 accepted, 33
5410                    // against 48.8 tok/s while the monitor called the draft
5411                    // alone "paying"). Count it as the losing streak in one.
5412                    spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
5413                    spec_mon.tokens = 0.0;
5414                    spec_mon.fails = 3;
5415                    spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
5416                    hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
5417                    next_pos += 1;
5418                    continue 'decode;
5419                }
5420                // ── Speculative: draft t+2, verify in a fused pair ──
5421                Some(m) if !graph_spec && generated + 1 < max_tokens => {
5422                    let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
5423                    drafted += 1;
5424                    let emb1 = self.embed_single(t_next);
5425                    let emb2 = self.embed_single(draft);
5426                    let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
5427
5428                    inference::rms_norm_into(
5429                        &h1,
5430                        &self.weights.final_norm,
5431                        self.rms_eps,
5432                        self.norm_style,
5433                        &mut self.ws.n1,
5434                    );
5435                    let mut logits1 = self.lm_head_forward(&self.ws.n1);
5436                    let t_after = sampler::sample_with_scratch_pool(
5437                        &logits1,
5438                        &self.sampler_config,
5439                        self.sampler_config.penalty_past(&all_ids, bounded_native),
5440                        &mut self.rng,
5441                        &mut self.sampler_scratch,
5442                        self.pool.as_deref(),
5443                    );
5444                    if self.confidence_on {
5445                        confidence.push(sampler::top1_prob_pool(
5446                            self.pool.as_deref(),
5447                            &mut self.sampler_scratch,
5448                            &logits1,
5449                            t_after,
5450                            calib_temp,
5451                        ));
5452                    }
5453                    attention::recycle_buf(&mut logits1);
5454                    if trace_on {
5455                        // Speculative decode is mutually exclusive with
5456                        // dynamic routing (router is None here) — no skill.
5457                        traces.push(TokenTrace {
5458                            t: generated,
5459                            token_id: t_after,
5460                            confidence: confidence.last().copied().unwrap_or(0.0),
5461                            active_skill: None,
5462                            recon: None,
5463                            switched: false,
5464                        });
5465                    }
5466                    let stop = !commit!(t_after);
5467
5468                    if t_after == draft {
5469                        accepted += 1;
5470                        self.commit_linear_scratch();
5471                        let _ = self.mtp_step(m, &h1, t_after, next_pos);
5472                        hidden = h2;
5473                        next_pos += 2;
5474                    } else {
5475                        // The draft lane is wrong: roll its KV entry back.
5476                        for layer in &mut self.kv_cache.layers {
5477                            layer.truncate_last(1);
5478                        }
5479                        if !stop {
5480                            let _ = self.mtp_step(m, &h1, t_after, next_pos);
5481                            hidden = self.forward_layers(
5482                                &self.embed_single(t_after),
5483                                next_pos + 1,
5484                                None,
5485                            );
5486                        }
5487                        next_pos += 2;
5488                    }
5489                    if stop {
5490                        break 'decode;
5491                    }
5492                }
5493                // ── Vanilla: forward the sampled token ──
5494                _ => {
5495                    // ── DeepSeek-V4 speculative decode (CMF_DSV4_SPEC=1):
5496                    // draft five on the card, verify batched, commit the
5497                    // accepted prefix. Greedy only; a rejected token's state
5498                    // is restored and replayed, so output equals the walk. ──
5499                    #[cfg(feature = "gpu")]
5500                    if Self::dsv4_spec_on() && self.dsv4.is_some() {
5501                        static SAID: std::sync::Once = std::sync::Once::new();
5502                        SAID.call_once(|| {
5503                            eprintln!(
5504                                "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
5505                                !self.dsv4_mtp.is_empty(),
5506                                task_mask.is_none(),
5507                                router.is_none(),
5508                                !trace_on,
5509                                self.sampler_config.temperature < 1e-6,
5510                                self.sampler_config.repetition_penalty == 1.0,
5511                            );
5512                        });
5513                    }
5514                    #[cfg(feature = "gpu")]
5515                    if Self::dsv4_spec_on()
5516                        && self.dsv4.is_some()
5517                        && !self.dsv4_mtp.is_empty()
5518                        && task_mask.is_none()
5519                        && router.is_none()
5520                        && !trace_on
5521                        && self.sampler_config.temperature < 1e-6
5522                        && self.sampler_config.repetition_penalty == 1.0
5523                        && generated + 1 < max_tokens
5524                        && all_ids.len() >= 2
5525                        && generated >= dsv4_spec_retry_at
5526                    {
5527                        let tip_token = all_ids[all_ids.len() - 2];
5528                        let drafted0 = drafted;
5529                        let round = self.dsv4_spec_step(
5530                            tip_token,
5531                            t_next,
5532                            next_pos,
5533                            max_tokens.saturating_sub(generated),
5534                            &mut drafted,
5535                            &mut accepted,
5536                        );
5537                        if drafted > drafted0 {
5538                            let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
5539                            if useful {
5540                                dsv4_spec_bad = 0;
5541                            } else {
5542                                dsv4_spec_bad += 1;
5543                                if dsv4_spec_bad >= 2 {
5544                                    dsv4_spec_bad = 0;
5545                                    dsv4_spec_retry_at = generated.saturating_add(32);
5546                                    tracing::info!(
5547                                        "dsv4: draft не окупился дважды — точный walk на 32 токена"
5548                                    );
5549                                }
5550                            }
5551                        }
5552                        if let Some((extra, n_pos)) = round {
5553                            next_pos = n_pos;
5554                            let mut stopped = false;
5555                            for &id in &extra {
5556                                if self.confidence_on {
5557                                    confidence.push(0.0);
5558                                }
5559                                if !commit!(id) {
5560                                    stopped = true;
5561                                    break;
5562                                }
5563                            }
5564                            if stopped {
5565                                break 'decode;
5566                            }
5567                            continue 'decode;
5568                        }
5569                    }
5570                    self.graph_want_logits = fuse_lm;
5571                    // Greedy burst (CMF_MULTISTEP, default 8, 1 = off): while
5572                    // nothing observes per-token state — pure argmax sampling,
5573                    // no router/trace/confidence/mask — decode k tokens per
5574                    // submit and commit them wholesale. The trailing normal
5575                    // forward leaves logits for the loop top, as always.
5576                    let mut t_fwd = t_next;
5577                    let pure_greedy = self.sampler_config.temperature < 1e-6
5578                        && self.sampler_config.repetition_penalty == 1.0
5579                        && self.sampler_config.suppress_tokens.is_empty();
5580                    // Off by default: at every k the burst measured at or
5581                    // below the plain path on this graph shape (k=1 loses
5582                    // the argmax dispatches vs a 1 MB readback, k>=8 loses
5583                    // inter-step drains vs the saved sync). Experimental.
5584                    let burst_k = std::env::var("CMF_MULTISTEP")
5585                        .ok()
5586                        .and_then(|v| v.parse::<usize>().ok())
5587                        .unwrap_or(0);
5588                    if pure_greedy
5589                        && burst_k >= 1
5590                        && fuse_lm
5591                        && task_mask.is_none()
5592                        && router.is_none()
5593                        && !trace_on
5594                        && !self.confidence_on
5595                    {
5596                        let mut stopped = false;
5597                        loop {
5598                            let room = max_tokens.saturating_sub(generated);
5599                            if room <= 2 {
5600                                break;
5601                            }
5602                            let k = burst_k.min(room - 1);
5603                            if k < 1 {
5604                                break;
5605                            }
5606                            let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
5607                                if self
5608                                    .graph_failed
5609                                    .swap(false, std::sync::atomic::Ordering::Relaxed)
5610                                {
5611                                    self.finish_generation(&mut mtp, &mut router, true);
5612                                    return Err(
5613                                        "GPU token graph failed during greedy burst".to_string()
5614                                    );
5615                                }
5616                                break;
5617                            };
5618                            next_pos += k;
5619                            for &id in &ids {
5620                                if !commit!(id) {
5621                                    stopped = true;
5622                                    break;
5623                                }
5624                            }
5625                            if stopped {
5626                                break;
5627                            }
5628                            t_fwd = *ids.last().unwrap();
5629                        }
5630                        if stopped {
5631                            break 'decode;
5632                        }
5633                    }
5634                    // Metal: keep the draft head's cache in step through
5635                    // the trial's plain phase and a paused speculation —
5636                    // the pair (hidden, t_fwd) at next_pos−1, the step the
5637                    // round's draft 0 would take. Without it the head's
5638                    // cache lagged the trunk by every plain token for the
5639                    // rest of the generation: the batched warm-up declined
5640                    // every later round and its rows went one by one (a
5641                    // whole MTP step per accepted token), and the drafts
5642                    // attended a context with those tokens missing.
5643                    #[cfg(target_os = "macos")]
5644                    if graph_spec
5645                        && spec_watchdog_off
5646                        && next_pos > 0
5647                        && self.mtp_graph_mode == Some(true)
5648                        && crate::gpu::q1_force()
5649                    {
5650                        if let Some(m) = mtp.as_mut() {
5651                            let _ = self.mtp_step_metal(m, &hidden, t_fwd, next_pos - 1, false);
5652                        }
5653                    }
5654                    hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
5655                    next_pos += 1;
5656                    // Dynamic routing: the forward updated φ; ask the
5657                    // router whether to switch skills before the next token.
5658                    if let Some(r) = &mut router {
5659                        let phi = self.dyn_phi_ema.clone();
5660                        let decision = r.step(&phi, generated);
5661                        if let Some(new_active) = decision {
5662                            let _ = self.set_active_skill(new_active);
5663                        }
5664                        // Backfill this token's coherence + switch flag from
5665                        // the just-run eval (freshest measured values).
5666                        if trace_on {
5667                            if let Some(last) = traces.last_mut() {
5668                                let e = r.last_best_e();
5669                                last.recon = e.is_finite().then_some(e);
5670                                last.switched = decision.is_some();
5671                            }
5672                        }
5673                    }
5674                }
5675            }
5676        }
5677
5678        let cancelled = finish_reason == "cancelled";
5679        // A generation during which the router switched weights holds no
5680        // state any single overlay would produce (each switch cleared the
5681        // cache mid-sequence), so it leaves no reuse key behind either.
5682        let dyn_switched = router.as_ref().is_some_and(|r| !r.switches.is_empty());
5683        if mimo_spec {
5684            if let Some(st) = self.mimo_mtp.as_ref() {
5685                let line = st.stats.line();
5686                tracing::info!("{line}");
5687                if std::env::var_os("CMF_MIMO_MTP_STATS").is_some() {
5688                    eprintln!("{line}");
5689                }
5690            }
5691        }
5692        self.finish_generation(&mut mtp, &mut router, cancelled);
5693
5694        let output_ids = &all_ids[input_ids.len()..];
5695        // Forwarded = prompt + all generated but the LAST sampled token
5696        // (emitted without being fed back). Exact only without MTP —
5697        // reuse is gated off when MTP is active.
5698        let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
5699        // A MiMo speculative round that stopped on an accepted draft (EOS,
5700        // cancel) leaves verify rows past the committed stream in the cache:
5701        // never offer that cache for reuse. Neither does a router that
5702        // switched weights mid-sequence (no single overlay produced it).
5703        let consumed = std::mem::take(&mut all_ids);
5704        if dyn_switched {
5705            self.clear_sequence_state();
5706        } else if cancelled || mimo_spec || prompt_rows.is_some() {
5707            self.clear_history();
5708        } else {
5709            self.record_consumed_prefix(&consumed[..forwarded.min(consumed.len())], reuse_from);
5710        }
5711        all_ids = consumed;
5712        let output_ids = &all_ids[input_ids.len()..];
5713        confidence.truncate(output_ids.len()); // guard against any overshoot
5714        traces.truncate(output_ids.len());
5715        Ok(GenerateResult {
5716            text: self.tokenizer.decode(output_ids),
5717            token_ids: output_ids.to_vec(),
5718            prompt_tokens: input_ids.len(),
5719            tokens_generated: generated,
5720            finish_reason,
5721            mtp_drafted: drafted,
5722            mtp_accepted: accepted,
5723            token_confidence: confidence,
5724            traces,
5725        })
5726    }
5727
5728    /// One MTP step: feed `(hidden_p, token_{p+1})` into the draft head,
5729    /// advance its KV cache at position `p`, return the drafted token
5730    /// for position `p+2`.
5731    fn mtp_step(
5732        &mut self,
5733        m: &mut MtpModule,
5734        hidden: &[f32],
5735        next_token: u32,
5736        position: usize,
5737    ) -> u32 {
5738        self.mtp_step_h(m, hidden, next_token, position).0
5739    }
5740
5741    /// Tally for `CMF_MTP_CHAIN_PROBE`: per depth, how often the CHAIN is
5742    /// still an exact prefix of the real continuation. Printed every 128
5743    /// depth-0 samples so a killed run still shows its table.
5744    fn chain_probe_note(depth: usize, prefix_ok: bool) {
5745        use std::sync::Mutex;
5746        static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
5747        let mut t = T.lock().unwrap();
5748        if t.len() <= depth {
5749            t.resize(depth + 1, (0, 0));
5750        }
5751        t[depth].0 += 1;
5752        t[depth].1 += prefix_ok as u64;
5753        if depth == 0 && t[0].0 % 128 == 0 {
5754            let line: Vec<String> = t
5755                .iter()
5756                .enumerate()
5757                .map(|(d, (n, k))| {
5758                    format!(
5759                        "d{}={:.0}%({n})",
5760                        d + 1,
5761                        100.0 * *k as f64 / (*n).max(1) as f64
5762                    )
5763                })
5764                .collect();
5765            eprintln!("mtp-chain: {}", line.join(" "));
5766        }
5767    }
5768
5769    /// `mtp_step` that also hands back the block's own output hidden — the
5770    /// state a CHAINED draft feeds the next step, the way a multi-token
5771    /// speculative round iterates the head on itself.
5772    /// One MTP block step from (trunk hidden, token): the head's LOGITS
5773    /// and the block's own hidden for chaining. The draft is argmax of the
5774    /// logits on the greedy path and a draw from their post-chain
5775    /// distribution on the sampling path.
5776    fn mtp_step_hl(
5777        &mut self,
5778        m: &mut MtpModule,
5779        hidden: &[f32],
5780        next_token: u32,
5781        position: usize,
5782    ) -> (Vec<f32>, Vec<f32>) {
5783        // The graph arm: the MTP block as a one-layer token graph with the
5784        // head fused — device attention over the block's own KV mirror,
5785        // one submit for block + head, hidden and logits back together.
5786        // Decided once per generation (see `mtp_graph_mode`).
5787        #[cfg(target_os = "macos")]
5788        if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
5789            if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
5790                self.mtp_graph_mode = Some(true);
5791                return r;
5792            }
5793            if self.mtp_graph_mode == Some(true) {
5794                tracing::error!("mtp Metal graph failed after admission");
5795                self.clear_sequence_state();
5796                self.graph_failed
5797                    .store(true, std::sync::atomic::Ordering::Relaxed);
5798                self.cancel
5799                    .store(true, std::sync::atomic::Ordering::Relaxed);
5800                return (Vec::new(), Vec::new());
5801            }
5802            self.mtp_graph_mode = Some(false);
5803        }
5804        #[cfg(feature = "gpu")]
5805        if self.mtp_graph_mode != Some(false) {
5806            if !self.mtp_graph_ok(m) {
5807                if self.mtp_graph_mode == Some(true) {
5808                    // A mirror was already admitted, so a capability change
5809                    // cannot safely switch this request to the stale CPU
5810                    // cache.  Keep the same terminal contract as a failed
5811                    // token graph.
5812                    tracing::error!("mtp graph became unavailable after admission");
5813                    self.clear_sequence_state();
5814                    self.graph_failed
5815                        .store(true, std::sync::atomic::Ordering::Relaxed);
5816                    self.cancel
5817                        .store(true, std::sync::atomic::Ordering::Relaxed);
5818                    return (Vec::new(), Vec::new());
5819                }
5820                self.mtp_graph_mode = Some(false);
5821            } else {
5822                if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
5823                    self.mtp_graph_mode = Some(true);
5824                    return r;
5825                }
5826                if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
5827                    // A token graph can have admitted a persistent MTP/GDN
5828                    // mirror before its readback failed.  The CPU MTP cache
5829                    // is not a valid continuation in that state; leave the
5830                    // flag set so the generation caller returns through its
5831                    // terminal error path instead of silently switching
5832                    // arithmetic.
5833                    return (Vec::new(), Vec::new());
5834                }
5835                // `mtp_graph_ok` was true, so a None here means a refusal or
5836                // failure after graph admission.  Do not fall through to a
5837                // CPU cache whose rows may lag the device mirror.
5838                tracing::error!("mtp graph failed or declined after admission");
5839                self.clear_sequence_state();
5840                self.graph_failed
5841                    .store(true, std::sync::atomic::Ordering::Relaxed);
5842                self.cancel
5843                    .store(true, std::sync::atomic::Ordering::Relaxed);
5844                return (Vec::new(), Vec::new());
5845            }
5846        }
5847        // fc concat order is [enorm(embed); hnorm(hidden)] — EMBEDDING
5848        // FIRST. Verified by the oracle (converter/mtp_oracle.py):
5849        // [emb;hid] → 45.8% acceptance, [hid;emb] → 0.00%.
5850        let e = self.embed_single(next_token);
5851        let mut cat = vec![0.0f32; 2 * self.hidden_size];
5852        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5853        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5854        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5855        let mut x = vec![0.0f32; self.hidden_size];
5856        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5857
5858        // One standard transformer block over the MTP's own cache.
5859        let lw = &m.layer;
5860        inference::rms_norm_into(
5861            &x,
5862            &lw.input_norm,
5863            self.rms_eps,
5864            self.norm_style,
5865            &mut self.ws.n1,
5866        );
5867        let attn = match &lw.attn {
5868            // MLA models carry no MTP head; this path cannot see them.
5869            AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
5870            AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
5871            AttnKind::Full {
5872                wq,
5873                wk,
5874                wv,
5875                wo,
5876                q_norm,
5877                k_norm,
5878                output_gate,
5879                softplus_gate,
5880                bias,
5881            } => {
5882                let mut cfg = self.attn_cfg(position);
5883                cfg.q_norm = q_norm.as_deref();
5884                cfg.k_norm = k_norm.as_deref();
5885                cfg.output_gate = *output_gate;
5886                cfg.softplus_gate = softplus_gate
5887                    .as_ref()
5888                    .map(|(gate, per_head)| (gate, *per_head));
5889                cfg.bias = bias
5890                    .as_ref()
5891                    .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
5892                attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
5893            }
5894            AttnKind::Linear(_)
5895            | AttnKind::LinearGdn(_)
5896            | AttnKind::ShortConv(_)
5897            | AttnKind::Bounded(_) => {
5898                unreachable!("MTP block is full attention")
5899            }
5900        };
5901        for (i, &a) in attn.iter().enumerate() {
5902            x[i] += a;
5903        }
5904        inference::rms_norm_into(
5905            &x,
5906            &lw.post_norm,
5907            self.rms_eps,
5908            self.norm_style,
5909            &mut self.ws.p1,
5910        );
5911        let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
5912        for (i, &f) in ffn.iter().enumerate() {
5913            x[i] += f;
5914        }
5915
5916        inference::rms_norm_into(
5917            &x,
5918            &m.final_norm,
5919            self.rms_eps,
5920            self.norm_style,
5921            &mut self.ws.n1,
5922        );
5923        let lg = self.lm_head_forward(&self.ws.n1);
5924        (lg, x)
5925    }
5926
5927    /// `mtp_step_hl` reduced to the greedy draft: argmax of the head.
5928    fn mtp_step_h(
5929        &mut self,
5930        m: &mut MtpModule,
5931        hidden: &[f32],
5932        next_token: u32,
5933        position: usize,
5934    ) -> (u32, Vec<f32>) {
5935        let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
5936        let draft = sampler::argmax(&lg);
5937        attention::recycle_buf(&mut lg);
5938        (draft, x)
5939    }
5940
5941    /// One speculative round for the trial: rounds 1..5 of a `Spec` phase
5942    /// advance it (the monitor already averaged this round); after five,
5943    /// the plain phase runs (once — a known plain rate decides at once);
5944    /// a decided speculation keeps re-checking the rule every round and
5945    /// stops after four losing rounds in a row.
5946    fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
5947        match trial {
5948            SpecTrial::Spec { t0, gen0, rounds } => {
5949                let rounds = rounds + 1;
5950                if rounds >= 5 {
5951                    if mon.plain_ms > 0.0 {
5952                        let keep = mon.pays();
5953                        mon.fails = 0;
5954                        tracing::info!(
5955                            "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5956                            mon.tokens,
5957                            mon.round_ms,
5958                            mon.plain_ms,
5959                            if keep { "speculating" } else { "plain" }
5960                        );
5961                        SpecTrial::Decided {
5962                            spec: keep,
5963                            recheck_at: if keep { usize::MAX } else { generated + 128 },
5964                        }
5965                    } else if mon.pays() {
5966                        // Metal: the rounds land enough tokens each that no
5967                        // plain measurement is needed — keep speculating,
5968                        // and re-check every round (a losing streak sends
5969                        // the loop to the plain phase, below).
5970                        mon.fails = 0;
5971                        tracing::info!(
5972                            "speculation trial: {:.2} tok/round in {:.1} ms — speculating (plain not timed)",
5973                            mon.tokens,
5974                            mon.round_ms,
5975                        );
5976                        SpecTrial::Decided {
5977                            spec: true,
5978                            recheck_at: usize::MAX,
5979                        }
5980                    } else {
5981                        SpecTrial::Plain {
5982                            t0: std::time::Instant::now(),
5983                            gen0: generated,
5984                        }
5985                    }
5986                } else {
5987                    SpecTrial::Spec { t0, gen0, rounds }
5988                }
5989            }
5990            SpecTrial::Decided { spec: true, .. } => {
5991                if mon.pays() {
5992                    mon.fails = 0;
5993                    trial
5994                } else {
5995                    mon.fails += 1;
5996                    if mon.fails >= 4 {
5997                        if mon.plain_ms <= 0.0 {
5998                            // Metal, plain never timed: four doubtful rounds
5999                            // buy the (bounded) plain measurement, and the
6000                            // exact rule decides from it.
6001                            tracing::info!(
6002                                "speculation doubtful: {:.2} tok/round in {:.1} ms — timing plain",
6003                                mon.tokens,
6004                                mon.round_ms,
6005                            );
6006                            return SpecTrial::Plain {
6007                                t0: std::time::Instant::now(),
6008                                gen0: generated,
6009                            };
6010                        }
6011                        tracing::info!(
6012                            "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
6013                            mon.tokens,
6014                            mon.round_ms,
6015                            mon.plain_ms
6016                        );
6017                        SpecTrial::Decided {
6018                            spec: false,
6019                            recheck_at: generated + 128,
6020                        }
6021                    } else {
6022                        trial
6023                    }
6024                }
6025            }
6026            other => other,
6027        }
6028    }
6029
6030    /// The MTP block's device-mirror id: the trunk's id with a high bit,
6031    /// so the (kv_id, layer) mirror keys never collide.
6032    fn mtp_kv_id(&self) -> u64 {
6033        self.graph_kv_id | (1u64 << 40)
6034    }
6035
6036    /// The MTP block's mirror layer index: 0 — its own kv_id keeps it
6037    /// apart from the trunk, and the BATCH graph (the warm-up path) keys
6038    /// its mirrors at layer 0 with no base of its own, so the draft's
6039    /// token graph must key the same slot.
6040    const MTP_LAYER_BASE: usize = 0;
6041
6042    /// The wgpu MTP draft writes speculative rows straight into its device
6043    /// mirror while the CPU owner retains only the real prompt/decode anchor.
6044    /// After verification, move that mirror cursor back to the anchor before
6045    /// replaying accepted pairs.  The next graph append then sees the same
6046    /// contiguous position as the CPU/Metal path without uploading stale
6047    /// speculative rows.
6048    #[cfg(feature = "gpu")]
6049    fn rewind_mtp_graph_mirror(&self, stored: usize) -> bool {
6050        self.mtp_graph_mode != Some(true)
6051            || crate::gpu::graph_kv_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, stored)
6052    }
6053
6054    /// A speculative verify graph appends the full `k+1` trunk rows before
6055    /// the acceptance count is known.  GDN state already has a snapshot
6056    /// restore; Full-attention mirrors need the matching logical cursor
6057    /// rewind so the next graph call does not reject an ahead-of-position KV
6058    /// cache after a partial acceptance.
6059    #[cfg(feature = "gpu")]
6060    fn rewind_trunk_graph_mirrors(&self, stored: usize) -> bool {
6061        let mut ok = true;
6062        let mut expected = false;
6063        for li in 0..self.num_layers {
6064            if matches!(
6065                self.weights.layers[self.phys_layer(li)].attn,
6066                AttnKind::Full { .. }
6067            ) {
6068                expected = true;
6069                ok &= crate::gpu::graph_kv_set_stored(self.graph_kv_id, li, stored);
6070            }
6071        }
6072        !expected || ok
6073    }
6074
6075    /// Count the recurrent layers participating in the trunk verify graph.
6076    /// Snapshot restore is all-or-nothing across that set; deriving the count
6077    /// from the model keeps the restore contract valid for looped models too.
6078    fn graph_gdn_layer_count(&self) -> usize {
6079        (0..self.num_layers)
6080            .filter(|&li| {
6081                matches!(
6082                    &self.weights.layers[self.phys_layer(li)].attn,
6083                    AttnKind::LinearGdn(_)
6084                )
6085            })
6086            .count()
6087    }
6088
6089    /// The block's input from (trunk hidden, token): eh_proj · [enorm(e);
6090    /// hnorm(h)] — the same arithmetic the per-op path starts with.
6091    fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
6092        let e = self.embed_single(next_token);
6093        let mut cat = vec![0.0f32; 2 * self.hidden_size];
6094        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6095        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6096        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6097        let mut x = vec![0.0f32; self.hidden_size];
6098        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6099        x
6100    }
6101
6102    /// Is the MTP block graphable at all (device up, full attention
6103    /// without softplus, dense FFN)? The plan itself is built per call.
6104    #[cfg(feature = "gpu")]
6105    fn mtp_block_graph_ok(&self, m: &MtpModule) -> bool {
6106        if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
6107            return false;
6108        }
6109        if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
6110            || !crate::gpu::enabled_here()
6111            || self.attn_softcap > 0.0
6112            || self.attention_heads_per_layer.is_some()
6113            // The block graph caches V as wide as K and feeds o_proj
6114            // nh·head_dim; a narrow-V model keeps its MTP block per-op.
6115            || self.v_head_dim.is_some()
6116        {
6117            return false;
6118        }
6119        matches!(
6120            &m.layer.attn,
6121            AttnKind::Full {
6122                softplus_gate: None,
6123                ..
6124            }
6125        ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
6126    }
6127
6128    /// Full MTP token-graph eligibility, including the fused lm-head and all
6129    /// block projection weights.  Keep this distinct from the block-only
6130    /// check: prompt warm-up does not need the head, while a draft step does.
6131    #[cfg(feature = "gpu")]
6132    fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
6133        if !self.mtp_block_graph_ok(m) {
6134            return false;
6135        }
6136        let AttnKind::Full { wq, wk, wv, wo, .. } = &m.layer.attn else {
6137            return false;
6138        };
6139        let FfnKind::Dense(d) = &m.layer.ffn else {
6140            return false;
6141        };
6142        d.segs.is_empty()
6143            && wq.graph_weight().is_some()
6144            && wk.graph_weight().is_some()
6145            && wv.graph_weight().is_some()
6146            && wo.graph_weight().is_some()
6147            && d.gate_proj.graph_weight().is_some()
6148            && d.up_proj.graph_weight().is_some()
6149            && d.down_proj.graph_weight().is_some()
6150            && self.weights.lm_head.graph_weight().is_some()
6151    }
6152
6153    /// One MTP block step on the wgpu token graph: block + fused head in
6154    /// one submit, the block hidden and the logits read back together.
6155    /// None = the graph cannot take this block (softplus gate, non-dense
6156    /// FFN, unquantized head, no device) — the caller keeps the per-op
6157    /// path for the whole generation.
6158    #[cfg(feature = "gpu")]
6159    fn mtp_step_graph(
6160        &mut self,
6161        m: &mut MtpModule,
6162        hidden: &[f32],
6163        next_token: u32,
6164        position: usize,
6165    ) -> Option<(Vec<f32>, Vec<f32>)> {
6166        if !self.mtp_graph_ok(m) {
6167            return None;
6168        }
6169        let lw = &m.layer;
6170        let AttnKind::Full {
6171            wq,
6172            wk,
6173            wv,
6174            wo,
6175            q_norm,
6176            k_norm,
6177            output_gate,
6178            softplus_gate,
6179            bias,
6180        } = &lw.attn
6181        else {
6182            return None;
6183        };
6184        if softplus_gate.is_some() {
6185            return None;
6186        }
6187        let FfnKind::Dense(d) = &lw.ffn else {
6188            return None;
6189        };
6190        if !d.segs.is_empty() {
6191            return None; // tube layers run on the segmented path
6192        }
6193        // The block's input first: it borrows `self` mutably (embed scratch,
6194        // pool), the plan below borrows the weights immutably.
6195        let mut x = self.mtp_block_input(m, hidden, next_token);
6196        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6197            let (_, i, kind, rs) = t.graph_weight()?;
6198            Some(crate::gpu::GraphW {
6199                idx: i,
6200                kind,
6201                row_scale: rs,
6202                data: &[],
6203                prism: crate::gpu::GraphPrismOp::None,
6204                affine: false,
6205            })
6206        }
6207        let (model, _, _, _) = wq.graph_weight()?;
6208        let model = model.clone();
6209        let (lm_gw, lm_rows) = {
6210            let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6211            // The draft's head over the CMF_DRAFT_VOCAB shortlist (the same
6212            // cut the native Metal draft takes): 662 MB a step on Qwen3.8
6213            // becomes 170 MB at 65536; the verify keeps the full head.
6214            let rows = if kind == 6 {
6215                self.draft_head_rows(self.weights.lm_head.rows())
6216            } else {
6217                self.weights.lm_head.rows()
6218            };
6219            (
6220                crate::gpu::GraphW {
6221                    idx: i,
6222                    kind,
6223                    row_scale: rs,
6224                    data: &[],
6225                    prism: crate::gpu::GraphPrismOp::None,
6226                    affine: false,
6227                },
6228                rows,
6229            )
6230        };
6231        let layer = crate::gpu::GraphLayer {
6232            input_norm: &lw.input_norm,
6233            attn: crate::gpu::GraphAttn::Full {
6234                wq: gw(wq)?,
6235                wk: gw(wk)?,
6236                wv: gw(wv)?,
6237                wo: gw(wo)?,
6238                q_norm: q_norm.as_deref(),
6239                k_norm: k_norm.as_deref(),
6240                late_qk_norm: self.qk_norm_after_rope,
6241                bias: bias
6242                    .as_ref()
6243                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6244                output_gate: *output_gate,
6245                cpu_k: m.kv.k_heads(),
6246                cpu_v: m.kv.v_heads(),
6247                geom: None,
6248                head_gate: None,
6249            },
6250            post_norm: &lw.post_norm,
6251            ffn: crate::gpu::GraphFfn::Dense {
6252                gate: gw(&d.gate_proj)?,
6253                up: gw(&d.up_proj)?,
6254                down: gw(&d.down_proj)?,
6255                act: d.act.graph_act()?,
6256            },
6257        };
6258        let nh = self.num_heads;
6259        let (nkv, hd, rd) = self.layer_geom(0);
6260        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6261        let mut logits = Vec::new();
6262        let ok = crate::gpu::forward_token_graph(
6263            &model,
6264            self.mtp_kv_id(),
6265            std::slice::from_ref(&layer),
6266            &[None],
6267            self.o1_epoch,
6268            &self.inv_freq,
6269            &mut x,
6270            nh,
6271            nkv,
6272            hd,
6273            self.attn_scale,
6274            rd,
6275            self.hidden_size,
6276            self.intermediate_size,
6277            position,
6278            self.kv_cache.max_seq_len,
6279            gemma,
6280            self.rms_eps as f32,
6281            Some((&lm_gw, lm_rows)),
6282            &m.final_norm,
6283            &mut logits,
6284            &[],
6285            1,
6286            None,
6287            None,
6288            None,
6289            Self::MTP_LAYER_BASE,
6290            true,
6291        );
6292        match ok {
6293            crate::gpu::TokenGraphOutcome::Completed => {}
6294            crate::gpu::TokenGraphOutcome::Declined => return None,
6295            crate::gpu::TokenGraphOutcome::Failed => {
6296                // The backend has already admitted persistent state.  Keep
6297                // this distinct from a capability refusal so the caller
6298                // cannot switch to the stale CPU MTP cache.
6299                self.clear_sequence_state();
6300                self.graph_failed
6301                    .store(true, std::sync::atomic::Ordering::Relaxed);
6302                self.cancel
6303                    .store(true, std::sync::atomic::Ordering::Relaxed);
6304                return None;
6305            }
6306        }
6307        logits.resize(self.vocab_size, 0.0);
6308        Some((logits, x))
6309    }
6310
6311    /// The warm-ups of one speculative round on the device: every accepted
6312    /// (hidden, token) pair as ONE batched graph run over the MTP block
6313    /// (no head) — its kv_append lands the pairs in the block's mirror.
6314    /// `pairs` are consecutive positions from `first_pos`.  The tri-state
6315    /// result is intentional: a refusal before admission may use the
6316    /// per-row/CPU route, while a failure after admission must terminate the
6317    /// sequence rather than fall through to a stale CPU cache.
6318    #[cfg(feature = "gpu")]
6319    fn mtp_warm_graph(
6320        &mut self,
6321        m: &mut MtpModule,
6322        pairs: &[(&[f32], u32)],
6323        first_pos: usize,
6324    ) -> crate::gpu::BatchGraphOutcome {
6325        if pairs.is_empty() {
6326            return crate::gpu::BatchGraphOutcome::Completed;
6327        }
6328        if !self.mtp_block_graph_ok(m) {
6329            return crate::gpu::BatchGraphOutcome::Declined;
6330        }
6331        let hs = self.hidden_size;
6332        // Block inputs for every pair (eh_proj on the per-op path, one
6333        // matvec each — the plan's own prologue).
6334        let mut hiddens = Vec::with_capacity(pairs.len() * hs);
6335        for (h, t) in pairs {
6336            hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
6337        }
6338        let lw = &m.layer;
6339        let AttnKind::Full {
6340            wq,
6341            wk,
6342            wv,
6343            wo,
6344            q_norm,
6345            k_norm,
6346            output_gate,
6347            bias,
6348            ..
6349        } = &lw.attn
6350        else {
6351            return crate::gpu::BatchGraphOutcome::Declined;
6352        };
6353        let FfnKind::Dense(d) = &lw.ffn else {
6354            return crate::gpu::BatchGraphOutcome::Declined;
6355        };
6356        if !d.segs.is_empty() {
6357            return crate::gpu::BatchGraphOutcome::Declined; // tube layers run on the segmented path
6358        }
6359        // The graph has no arm for a projected output gate, and only the
6360        // activations `graph_act` names.
6361        let (AttnKind::Full { softplus_gate: None, .. }, Some(gact)) =
6362            (&lw.attn, d.act.graph_act())
6363        else {
6364            return crate::gpu::BatchGraphOutcome::Declined;
6365        };
6366        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6367            let (_, i, kind, rs) = t.graph_weight()?;
6368            Some(crate::gpu::GraphW {
6369                idx: i,
6370                kind,
6371                row_scale: rs,
6372                data: &[],
6373                prism: crate::gpu::GraphPrismOp::None,
6374                affine: false,
6375            })
6376        }
6377        let Some((model, _, _, _)) = wq.graph_weight() else {
6378            return crate::gpu::BatchGraphOutcome::Declined;
6379        };
6380        let model = model.clone();
6381        let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
6382            gw(wq),
6383            gw(wk),
6384            gw(wv),
6385            gw(wo),
6386            gw(&d.gate_proj),
6387            gw(&d.up_proj),
6388            gw(&d.down_proj),
6389        ) else {
6390            return crate::gpu::BatchGraphOutcome::Declined;
6391        };
6392        let layer = crate::gpu::GraphLayer {
6393            input_norm: &lw.input_norm,
6394            attn: crate::gpu::GraphAttn::Full {
6395                wq: gwq,
6396                wk: gwk,
6397                wv: gwv,
6398                wo: gwo,
6399                q_norm: q_norm.as_deref(),
6400                k_norm: k_norm.as_deref(),
6401                late_qk_norm: self.qk_norm_after_rope,
6402                bias: bias
6403                    .as_ref()
6404                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6405                output_gate: *output_gate,
6406                cpu_k: m.kv.k_heads(),
6407                cpu_v: m.kv.v_heads(),
6408                geom: None,
6409                head_gate: None,
6410            },
6411            post_norm: &lw.post_norm,
6412            ffn: crate::gpu::GraphFfn::Dense {
6413                gate: gg,
6414                up: gu,
6415                down: gd,
6416                act: gact,
6417            },
6418        };
6419        let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
6420        let nh = self.num_heads;
6421        let (nkv, hd, rd) = self.layer_geom(0);
6422        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6423        crate::gpu::forward_batch_graph(
6424            &model,
6425            self.mtp_kv_id(),
6426            std::slice::from_ref(&layer),
6427            &self.inv_freq,
6428            &mut hiddens,
6429            nh,
6430            nkv,
6431            hd,
6432            rd,
6433            hs,
6434            self.intermediate_size,
6435            &positions,
6436            self.kv_cache.max_seq_len,
6437            gemma,
6438            self.rms_eps as f32,
6439            self.attn_scale,
6440            pairs.len(),
6441            &[],
6442            0,
6443            None,
6444            None,
6445        )
6446    }
6447
6448    /// Complete an MTP warm-up after the batched graph has refused.  A
6449    /// graphable block is retried one row at a time; once any device row has
6450    /// been admitted, a CPU fallback would observe a stale mirror, so every
6451    /// token-graph refusal is terminal.  If the block is not graphable and no
6452    /// mirror exists yet, warming on the CPU is safe and records the CPU mode
6453    /// for the rest of the generation.
6454    #[cfg(feature = "gpu")]
6455    fn mtp_warm_graph_fallback(
6456        &mut self,
6457        m: &mut MtpModule,
6458        pairs: &[(&[f32], u32)],
6459        first_pos: usize,
6460    ) -> bool {
6461        if pairs.is_empty() {
6462            return true;
6463        }
6464        let graphable = self.mtp_block_graph_ok(m);
6465        if !graphable {
6466            // A previously admitted mirror cannot be made coherent by
6467            // appending to the host cache.  The caller turns this into a
6468            // terminal generation error and clears both mirrors.
6469            if self.mtp_graph_mode == Some(true) {
6470                return false;
6471            }
6472            self.mtp_graph_mode = Some(false);
6473            for (j, (h, t)) in pairs.iter().enumerate() {
6474                self.mtp_warm(m, h, *t, first_pos + j);
6475            }
6476            return true;
6477        }
6478
6479        // The batch refusal is recoverable only through the same device
6480        // state.  Keep rows owned until each token graph has completed; a
6481        // None is treated as unsafe because the token-graph API deliberately
6482        // collapses its backend refusal/failure into that result.
6483        for (j, (h, t)) in pairs.iter().enumerate() {
6484            if self.mtp_step_graph(m, h, *t, first_pos + j).is_none() {
6485                return false;
6486            }
6487        }
6488        self.mtp_graph_mode = Some(true);
6489        true
6490    }
6491
6492    /// Warm a contiguous set of MTP pairs using the existing graph seam, with
6493    /// an all-or-nothing error contract for callers that already admitted the
6494    /// trunk batch.  The non-GPU build keeps the same pair accounting while
6495    /// using the established CPU warm path.
6496    #[cfg(feature = "gpu")]
6497    fn mtp_warm_prefill_pairs(
6498        &mut self,
6499        m: &mut MtpModule,
6500        pairs: &[(&[f32], u32)],
6501        first_pos: usize,
6502    ) -> Result<(), &'static str> {
6503        // Keep unsupported token-graph heads on the established CPU MTP
6504        // route before admitting any block mirror.  Once a device mirror is
6505        // active, the same condition is terminal because CPU rows cannot
6506        // repair its state.
6507        if self.mtp_graph_mode == Some(false) || !self.mtp_graph_ok(m) {
6508            if self.mtp_graph_mode == Some(true) {
6509                return Err("MTP token graph became unavailable after admission");
6510            }
6511            self.mtp_graph_mode = Some(false);
6512            for (j, (h, t)) in pairs.iter().enumerate() {
6513                self.mtp_warm(m, h, *t, first_pos + j);
6514            }
6515            return Ok(());
6516        }
6517        match self.mtp_warm_graph(m, pairs, first_pos) {
6518            crate::gpu::BatchGraphOutcome::Completed => {
6519                if !pairs.is_empty() {
6520                    self.mtp_graph_mode = Some(true);
6521                }
6522                Ok(())
6523            }
6524            crate::gpu::BatchGraphOutcome::Declined => {
6525                if self.mtp_warm_graph_fallback(m, pairs, first_pos) {
6526                    Ok(())
6527                } else {
6528                    Err("MTP warm-up fallback failed after device admission")
6529                }
6530            }
6531            crate::gpu::BatchGraphOutcome::Failed => {
6532                Err("MTP warm batch graph failed after admission")
6533            }
6534        }
6535    }
6536
6537    #[cfg(not(feature = "gpu"))]
6538    fn mtp_warm_prefill_pairs(
6539        &mut self,
6540        m: &mut MtpModule,
6541        pairs: &[(&[f32], u32)],
6542        first_pos: usize,
6543    ) -> Result<(), &'static str> {
6544        for (j, (h, t)) in pairs.iter().enumerate() {
6545            self.mtp_warm(m, h, *t, first_pos + j);
6546        }
6547        Ok(())
6548    }
6549
6550    /// The MTP block alone — advance its KV with a (hidden, token) pair the
6551    /// verify just proved, without paying the head. What keeps the draft's
6552    /// attention context warm between speculative rounds.
6553    fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
6554        let e = self.embed_single(next_token);
6555        let mut cat = vec![0.0f32; 2 * self.hidden_size];
6556        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6557        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6558        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6559        let mut x = vec![0.0f32; self.hidden_size];
6560        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6561        inference::rms_norm_into(
6562            &x,
6563            &m.layer.input_norm,
6564            self.rms_eps,
6565            self.norm_style,
6566            &mut self.ws.n1,
6567        );
6568        let attn = match &m.layer.attn {
6569            AttnKind::Full {
6570                wq,
6571                wk,
6572                wv,
6573                wo,
6574                q_norm,
6575                k_norm,
6576                output_gate,
6577                softplus_gate,
6578                bias,
6579            } => {
6580                let mut cfg = self.attn_cfg(position);
6581                cfg.q_norm = q_norm.as_deref();
6582                cfg.k_norm = k_norm.as_deref();
6583                cfg.output_gate = *output_gate;
6584                cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
6585                cfg.bias = bias
6586                    .as_ref()
6587                    .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
6588                attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
6589            }
6590            _ => return,
6591        };
6592        let _ = attn;
6593    }
6594
6595    /// Speculative decode ON the wgpu whole-token graph: draft k with the
6596    /// MTP head, verify all of them plus the tip in ONE batched graph
6597    /// submit whose tail folds the head, commit the accepted prefix and
6598    /// roll the GDN state back to the last real position. Greedy only —
6599    /// output equals the plain graph's token for token, the way the DSV4
6600    /// verify equals the walk.
6601    #[cfg(feature = "gpu")]
6602    #[allow(clippy::too_many_arguments)]
6603    fn graph_spec_step(
6604        &mut self,
6605        m: &mut MtpModule,
6606        hidden: &[f32],
6607        t_next: u32,
6608        next_pos: usize,
6609        drafted: &mut usize,
6610        accepted: &mut usize,
6611        // The committed stream (prompt + generated so far, `t_next`
6612        // included): the sampler chain's penalties read it, and the
6613        // sampling arm extends it with the drafts position by position.
6614        all_ids: &mut Vec<u32>,
6615        // Tokens left before `max_tokens`. A round commits up to k
6616        // accepted drafts, and those positions are already in the cache,
6617        // so the depth is capped here — trimming the output afterwards
6618        // would leave cache rows the committed stream does not have.
6619        room: usize,
6620    ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
6621        // 3 is the measured optimum on Qwen3.6-27B / RTX 5090 (medians
6622        // of three, greedy): 51.1 tok/s against a plain 49.4, where k=2
6623        // gives 46.1, k=4 50.0, k=5 47.4, k=6 45.2. Acceptance is 89-91%
6624        // throughout — what turns the curve over is the verify, which
6625        // costs ~7.4 ms per extra position, and the draft ~3 ms a step.
6626        // 4 since the draft moved onto the graph (Qwen3.8-27B / 5090:
6627        // k=3 51.2, k=4 51.8 with the per-op draft; the graph draft
6628        // halves the draft cost, so the extra draft is cheaper still).
6629        // 5 with the int8 verify (the default: measured 76.5 against
6630        // k=4's 72-74 and k=6's 74 on the 5090), 4 with the f32 one.
6631        #[cfg(target_os = "macos")]
6632        let metal_native = crate::gpu::q1_force();
6633        #[cfg(not(target_os = "macos"))]
6634        let metal_native = false;
6635        #[cfg(feature = "gpu")]
6636        let k_default = if metal_native {
6637            // the Metal verify's GEMM tile is 8 rows wide and flat in b:
6638            // seven drafts + the tip fill it for free
6639            7
6640        } else if crate::gpu_wgpu::verify_i8_on() {
6641            5
6642        } else {
6643            4
6644        };
6645        #[cfg(not(feature = "gpu"))]
6646        let k_default = 4;
6647        let k_env: Option<usize> = std::env::var("CMF_GRAPH_SPEC_K")
6648            .ok()
6649            .and_then(|v| v.parse().ok())
6650            .filter(|&v| (1..=8).contains(&v));
6651        // Adaptive depth: start below the card's flat-verify optimum and
6652        // let the accepted fraction move it — predictable text climbs to
6653        // the old default within a few rounds, prose settles at 2-3 where
6654        // the shorter verify pays.
6655        let (k_start, k_max) = if metal_native { (7, 7) } else { (3, k_default.max(5)) };
6656        let k_full: usize = k_env.unwrap_or_else(|| self.spec_k_adapt.unwrap_or(k_start));
6657        let k_spec = k_full.min(room).max(1);
6658        // a tail round cut short by `room` says nothing about the text:
6659        // it must not move the adaptive depth the next request starts at
6660        let k_capped = k_spec < k_full;
6661        if next_pos == 0 {
6662            return None;
6663        }
6664        let t_round = std::time::Instant::now();
6665        // Submissions per phase — and they say where the round's money is.
6666        // Qwen3.6-27B on an RTX 5090, k=3:
6667        //
6668        //   draft   9.3 ms / 12 submissions   (four per MTP step)
6669        //   verify 52.8 ms /  1               (the batched graph)
6670        //   commit  5.4 ms /  6               (two per warm)
6671        //
6672        // The verify is already one submit. The draft's own work is 834 MB
6673        // a step — 0.8 ms at this card's measured 1056 GB/s — against 3.1
6674        // ms measured, so ~0.58 ms of every step is round trip, not
6675        // arithmetic, and the same holds for the warms. Eighteen round
6676        // trips a round at roughly half a millisecond each is ~11 ms of a
6677        // 68 ms round: fusing the MTP block into ONE submit the way the
6678        // trunk already is projects to ~64 tok/s against today's 50.9.
6679        // That is the largest measured item left on this path.
6680        let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
6681        let sub0 = subs();
6682        // Greedy without penalties verifies by argmax equality (bit-exact
6683        // against the plain path). Anything else is speculative SAMPLING:
6684        // each draft is a DRAW from the MTP head's post-chain distribution
6685        // q_j, kept for the accept test; the verify's rows give p_j.
6686        let cfg = self.sampler_config.clone();
6687        let penalized = !(cfg.repetition_penalty == 1.0
6688            && cfg.presence_penalty == 0.0
6689            && cfg.suppress_tokens.is_empty());
6690        // Three verify regimes: plain greedy (argmax of the raw rows),
6691        // greedy WITH penalties (argmax of the penalized rows — a single
6692        // pass each, no distributions), and sampling (draw / accept /
6693        // correct on post-chain distributions).
6694        let greedy_pen = cfg.temperature < 1e-6 && penalized;
6695        let sampling = cfg.temperature >= 1e-6;
6696        // Sampling with a top-k goes through the SPARSE chain: the dense
6697        // one builds nine 248k-float distributions a round (four drafts,
6698        // five verify rows) and measured 19-22 tok/s against a plain 40 —
6699        // the host, not the card. Sparse, the same nine cost tens of
6700        // microseconds each.
6701        let sparse = sampling && sampler::sparse_ok(&cfg);
6702        let base_len = all_ids.len();
6703        if sampling && !sparse && self.spec_q.len() < k_spec {
6704            self.spec_q.resize_with(k_spec, Vec::new);
6705        }
6706        if sparse && self.spec_qs.len() < k_spec {
6707            self.spec_qs.resize_with(k_spec, Vec::new);
6708        }
6709        // Draft the chain: first from the trunk's tip hidden, then the head
6710        // iterating on itself. Rows land in the MTP KV; the chain rows past
6711        // the first are speculation over speculative state and roll back
6712        // below, replaced by verified pairs.
6713        let mut drafts = Vec::with_capacity(k_spec);
6714        let mut hx = hidden.to_vec();
6715        // CMF_SPEC_DBG=1: draft 0 through BOTH MTP arms (graph and per-op)
6716        // from the same inputs — are the arms the difference, or the inputs?
6717        let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
6718        spec_stamp("pro");
6719        // Plain greedy on native Metal: the whole chain as one command
6720        // buffer (device argmax + embedding gather between the steps).
6721        // A decline before commit hands the round to the per-step loop
6722        // below; a failure after commit is terminal, like any graph
6723        // failure after admission.
6724        #[cfg(target_os = "macos")]
6725        if metal_native && !sampling && !greedy_pen && self.mtp_graph_mode != Some(false) {
6726            match self.mtp_draft_chain_metal(m, hidden, t_next, next_pos - 1, k_spec) {
6727                Ok(ids) => {
6728                    self.mtp_graph_mode = Some(true);
6729                    drafts = ids;
6730                }
6731                Err(true) => {
6732                    tracing::error!("mtp Metal draft chain failed after commit");
6733                    self.clear_sequence_state();
6734                    self.graph_failed
6735                        .store(true, std::sync::atomic::Ordering::Relaxed);
6736                    self.cancel
6737                        .store(true, std::sync::atomic::Ordering::Relaxed);
6738                    return None;
6739                }
6740                Err(false) => {}
6741            }
6742        }
6743        for j in drafts.len()..k_spec {
6744            let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
6745            let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
6746            if spec_dbg {
6747                let saved = self.mtp_graph_mode;
6748                self.mtp_graph_mode = Some(false);
6749                let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6750                self.mtp_graph_mode = saved;
6751                if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6752                    return None;
6753                }
6754                m.kv.truncate_last(1);
6755                dbg_ref = Some(r);
6756            }
6757            let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6758            if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6759                return None;
6760            }
6761            if let Some((lg_cpu, h_cpu)) = dbg_ref {
6762                let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
6763                let dl = lg
6764                    .iter()
6765                    .zip(&lg_cpu)
6766                    .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6767                let dh = hj
6768                    .iter()
6769                    .zip(&h_cpu)
6770                    .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6771                eprintln!(
6772                    "spec-dbg j={j} pos {} tok_in {tok_in}: per-op draft {} graph draft {} | max|dlogit| {dl:.3} | |h_cpu| {:.2} |h_graph| {:.2} max|dh| {dh:.3} | kv rows {}",
6773                    next_pos - 1 + j,
6774                    sampler::argmax(&lg_cpu),
6775                    sampler::argmax(&lg),
6776                    n(&h_cpu),
6777                    n(&hj),
6778                    m.kv.seq_len
6779                );
6780            }
6781            let dj = if sparse {
6782                let mut q = std::mem::take(&mut self.spec_qs[j]);
6783                let ok = sampler::sparse_distribution_into(
6784                    &lg,
6785                    &cfg,
6786                    all_ids,
6787                    &mut self.sampler_scratch,
6788                    self.pool.as_deref(),
6789                    &mut q,
6790                );
6791                let d = if ok {
6792                    sampler::draw_sparse(&q, &mut self.rng)
6793                } else {
6794                    // everything filtered: the dense chain's greedy fallback
6795                    let t = sampler::argmax(&lg);
6796                    q.clear();
6797                    q.push((t, 1.0));
6798                    t
6799                };
6800                self.spec_qs[j] = q;
6801                all_ids.push(d);
6802                d
6803            } else if sampling {
6804                let mut q = std::mem::take(&mut self.spec_q[j]);
6805                sampler::distribution_into(
6806                    &lg,
6807                    &cfg,
6808                    all_ids,
6809                    &mut self.sampler_scratch,
6810                    self.pool.as_deref(),
6811                    &mut q,
6812                );
6813                let d = sampler::draw(&q, &mut self.rng);
6814                self.spec_q[j] = q;
6815                all_ids.push(d); // the next draft's penalties see this one
6816                d
6817            } else if greedy_pen {
6818                let d = sampler::argmax_penalized(
6819                    &lg,
6820                    &cfg,
6821                    all_ids,
6822                    &mut self.sampler_scratch,
6823                    self.pool.as_deref(),
6824                );
6825                all_ids.push(d);
6826                d
6827            } else {
6828                sampler::argmax(&lg)
6829            };
6830            attention::recycle_buf(&mut lg);
6831            drafts.push(dj);
6832            hx = hj;
6833            spec_stamp("d.pick");
6834        }
6835        all_ids.truncate(base_len);
6836        *drafted += k_spec;
6837        let t_draft = t_round.elapsed();
6838        let sub_draft = subs();
6839        // Verify batch: [t_next, d1 .. d_{k-1}] at next_pos.. — every row's
6840        // logits come back from the graph's own head.
6841        let b = k_spec + 1;
6842        let mut hiddens = vec![0.0f32; b * self.hidden_size];
6843        for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
6844            let e = self.embed_single(t);
6845            hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
6846        }
6847        let positions: Vec<usize> = (next_pos..next_pos + b).collect();
6848        spec_stamp("v.emb");
6849        let (lm_gw, lm_rows) = {
6850            let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6851            (
6852                crate::gpu::GraphW {
6853                    idx: i,
6854                    kind,
6855                    row_scale: rs,
6856                    data: &[],
6857                    prism: crate::gpu::GraphPrismOp::None,
6858                    affine: false,
6859                },
6860                self.weights.lm_head.rows(),
6861            )
6862        };
6863        let mut logits = Vec::new();
6864        let final_norm = self.weights.final_norm.clone();
6865        // Plain greedy on Metal: the b argmaxes come from the device
6866        // (`argmax_rows` after the head) and the 7.9 MB logits plane is
6867        // never read back — the round's decision needs only the ids, and
6868        // the loop top takes the last verified id as `spec_forced`, which
6869        // is exactly what its argmax of the row would give. The full rows
6870        // stay for anything that reads them: sampling, penalties,
6871        // confidence, the verify oracle, the logit dump.
6872        // `CMF_METAL_DEV_ARGMAX=0` keeps the host path.
6873        #[cfg(target_os = "macos")]
6874        let greedy_dev = metal_native
6875            && !sampling
6876            && !greedy_pen
6877            && !self.confidence_on
6878            && self.final_softcap.is_none()
6879            // The host acceptance argmax scans the WHOLE head row
6880            // (`lm_rows`), the sampler's own row only `vocab_size`: they
6881            // coincide exactly when the head has no padding rows, and
6882            // only then is the device argmax (which scores `vocab_size`)
6883            // bit-identical to both.
6884            && self.vocab_size == lm_rows
6885            && std::env::var_os("CMF_METAL_VERIFY_CHECK").is_none()
6886            && std::env::var_os("CMF_LOGIT_DUMP").is_none()
6887            && std::env::var("CMF_METAL_DEV_ARGMAX").as_deref() != Ok("0");
6888        #[cfg(not(target_os = "macos"))]
6889        let greedy_dev = false;
6890        let mut dev_ids: Vec<u32> = Vec::new();
6891        #[cfg(target_os = "macos")]
6892        let verify_outcome = if metal_native {
6893            let lm = self.weights.lm_head.q1_parts()?;
6894            let n_score = self.vocab_size.min(lm_rows);
6895            self.try_batch_graph_metal(
6896                &mut hiddens,
6897                &positions,
6898                b,
6899                Some((lm, &final_norm, &mut logits)),
6900                if greedy_dev {
6901                    Some((n_score, &mut dev_ids))
6902                } else {
6903                    None
6904                },
6905            )
6906        } else {
6907            self.try_batch_graph_wgpu(
6908                &mut hiddens,
6909                &positions,
6910                b,
6911                Some(crate::gpu::SpecTail {
6912                    lm: lm_gw,
6913                    lm_rows,
6914                    final_norm: &final_norm,
6915                    logits_out: &mut logits,
6916                }),
6917            )
6918        };
6919        #[cfg(not(target_os = "macos"))]
6920        let verify_outcome = self.try_batch_graph_wgpu(
6921            &mut hiddens,
6922            &positions,
6923            b,
6924            Some(crate::gpu::SpecTail {
6925                lm: lm_gw,
6926                lm_rows,
6927                final_norm: &final_norm,
6928                logits_out: &mut logits,
6929            }),
6930        );
6931        match verify_outcome {
6932            crate::gpu::BatchGraphOutcome::Completed => {}
6933            crate::gpu::BatchGraphOutcome::Declined => {
6934                // The verifier refused before admission.  Its draft MTP
6935                // rows are still device-resident, so rewind the separate
6936                // mirror before the caller takes the exact one-token path.
6937                m.kv.truncate_last(k_spec);
6938                if !metal_native && !self.rewind_mtp_graph_mirror(next_pos) {
6939                    self.clear_sequence_state();
6940                    self.graph_failed
6941                        .store(true, std::sync::atomic::Ordering::Relaxed);
6942                    self.cancel
6943                        .store(true, std::sync::atomic::Ordering::Relaxed);
6944                    tracing::error!("MTP graph mirror rewind failed after verify decline");
6945                }
6946                return None;
6947            }
6948            crate::gpu::BatchGraphOutcome::Failed => {
6949                // A failed batch may have advanced trunk/GDN state.  Clear
6950                // both mirrors and preserve the terminal outcome rather than
6951                // falling through to stale CPU state.
6952                self.clear_sequence_state();
6953                self.graph_failed
6954                    .store(true, std::sync::atomic::Ordering::Relaxed);
6955                self.cancel
6956                    .store(true, std::sync::atomic::Ordering::Relaxed);
6957                tracing::error!("MTP verify batch graph failed after admission");
6958                return None;
6959            }
6960        }
6961        // `CMF_METAL_VERIFY_CHECK=1`: run the same b tokens through the
6962        // plain per-token path and compare each row's argmax + logits with
6963        // the verify's — the bring-up oracle for the batched graph. The
6964        // plain forwards mutate the CPU state; it is snapshotted and put
6965        // back, and the K/V mirrors re-pointed, before the round goes on.
6966        #[cfg(target_os = "macos")]
6967        if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
6968            let snap: Vec<Vec<f32>> = self
6969                .kv_cache
6970                .layers
6971                .iter()
6972                .map(|l| l.linear_state.clone())
6973                .collect();
6974            let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
6975            let toks: Vec<u32> = std::iter::once(t_next)
6976                .chain(drafts.iter().copied())
6977                .collect();
6978            let want_save = self.graph_want_logits;
6979            self.graph_want_logits = false;
6980            for (i, &t) in toks.iter().enumerate() {
6981                let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
6982                let _ = self.graph_logits.take();
6983                // CMF_SPEC_PLAIN_HIDDEN=1: the next round drafts from the
6984                // plain path's hidden instead of the verify's (an experiment
6985                // on the chain's sensitivity to the half-GEMM noise)
6986                if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
6987                    hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
6988                }
6989                let ref_lg = self.logits_from_hidden(&hi);
6990                let row = &logits[i * lm_rows..(i + 1) * lm_rows];
6991                let ra = sampler::argmax(&ref_lg);
6992                let va = sampler::argmax(row);
6993                let mut md = 0f32;
6994                let mut rms = 0f64;
6995                for j in 0..lm_rows.min(ref_lg.len()) {
6996                    let d = (ref_lg[j] - row[j]).abs();
6997                    md = md.max(d);
6998                    rms += (d as f64) * (d as f64);
6999                }
7000                let mut hd = 0f32;
7001                for j in 0..self.hidden_size {
7002                    hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
7003                }
7004                eprintln!(
7005                    "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
7006                    next_pos + i,
7007                    if ra == va { "OK" } else { "MISMATCH" },
7008                    (rms / lm_rows as f64).sqrt()
7009                );
7010            }
7011            self.graph_want_logits = want_save;
7012            // restore IN PLACE: the pending verify graph wraps these very
7013            // allocations (zero-copy) — replacing the Vec would strand it
7014            for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7015                if l.linear_state.len() == st.len() {
7016                    l.linear_state.copy_from_slice(&st);
7017                } else {
7018                    l.linear_state = st;
7019                }
7020            }
7021            for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
7022                let extra = l.seq_len.saturating_sub(n0);
7023                if extra > 0 {
7024                    l.truncate_last(extra);
7025                    crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
7026                }
7027            }
7028        }
7029        let t_verify = t_round.elapsed();
7030        let sub_verify = subs();
7031        // Acceptance. Greedy: row i's argmax is the trunk's token after
7032        // input i. Sampling: accept draft i with min(1, p_i/q_i), and on
7033        // the first rejection draw the correction from max(0, p_i − q_i)
7034        // — that token is committed by the loop top as-is (spec_forced).
7035        let mut a = 0usize;
7036        let mut forced: Option<u32> = None;
7037        let ids: Vec<u32> = if sparse {
7038            let mut p = std::mem::take(&mut self.spec_ps);
7039            let mut res = std::mem::take(&mut self.spec_ress);
7040            while a < k_spec {
7041                let ok = sampler::sparse_distribution_into(
7042                    &logits[a * lm_rows..(a + 1) * lm_rows],
7043                    &cfg,
7044                    all_ids,
7045                    &mut self.sampler_scratch,
7046                    self.pool.as_deref(),
7047                    &mut p,
7048                );
7049                if !ok {
7050                    let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
7051                    p.clear();
7052                    p.push((t, 1.0));
7053                }
7054                match sampler::spec_accept_or_correct_sparse(
7055                    &p,
7056                    &self.spec_qs[a],
7057                    drafts[a],
7058                    &mut self.rng,
7059                    &mut res,
7060                ) {
7061                    None => {
7062                        all_ids.push(drafts[a]);
7063                        a += 1;
7064                    }
7065                    Some(c) => {
7066                        forced = Some(c);
7067                        break;
7068                    }
7069                }
7070            }
7071            all_ids.truncate(base_len);
7072            self.spec_ps = p;
7073            self.spec_ress = res;
7074            drafts.clone()
7075        } else if sampling {
7076            let mut p = std::mem::take(&mut self.spec_p);
7077            let mut res = std::mem::take(&mut self.spec_res);
7078            while a < k_spec {
7079                sampler::distribution_into(
7080                    &logits[a * lm_rows..(a + 1) * lm_rows],
7081                    &cfg,
7082                    all_ids,
7083                    &mut self.sampler_scratch,
7084                    self.pool.as_deref(),
7085                    &mut p,
7086                );
7087                match sampler::spec_accept_or_correct(
7088                    &p,
7089                    &self.spec_q[a],
7090                    drafts[a],
7091                    &mut self.rng,
7092                    &mut res,
7093                    self.pool.as_deref(),
7094                ) {
7095                    None => {
7096                        all_ids.push(drafts[a]);
7097                        a += 1;
7098                    }
7099                    Some(c) => {
7100                        forced = Some(c);
7101                        break;
7102                    }
7103                }
7104            }
7105            all_ids.truncate(base_len);
7106            self.spec_p = p;
7107            self.spec_res = res;
7108            // the accepted drafts ARE the verified tokens after inputs 0..a
7109            drafts.clone()
7110        } else if greedy_pen {
7111            // Row i's penalized argmax, penalties over the stream that
7112            // includes the accepted drafts before it — the plain loop's
7113            // exact arithmetic, one pass per row, no working copy.
7114            let mut ids: Vec<u32> = Vec::with_capacity(b);
7115            for i in 0..b {
7116                let t = sampler::argmax_penalized(
7117                    &logits[i * lm_rows..(i + 1) * lm_rows],
7118                    &cfg,
7119                    all_ids,
7120                    &mut self.sampler_scratch,
7121                    self.pool.as_deref(),
7122                );
7123                ids.push(t);
7124                if i < k_spec && t == drafts[i] {
7125                    all_ids.push(t);
7126                } else {
7127                    break;
7128                }
7129            }
7130            all_ids.truncate(base_len);
7131            while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
7132                a += 1;
7133            }
7134            // rows past the first mismatch were never scored; the loop
7135            // top re-samples the last verified row itself.
7136            ids
7137        } else if greedy_dev && dev_ids.len() == b {
7138            let ids = std::mem::take(&mut dev_ids);
7139            while a < k_spec && ids[a] == drafts[a] {
7140                a += 1;
7141            }
7142            ids
7143        } else {
7144            if logits.len() < b * lm_rows {
7145                // the device argmax was asked for and came back short:
7146                // no rows to fall back on — terminal like a failed batch
7147                self.clear_sequence_state();
7148                self.graph_failed
7149                    .store(true, std::sync::atomic::Ordering::Relaxed);
7150                self.cancel
7151                    .store(true, std::sync::atomic::Ordering::Relaxed);
7152                tracing::error!("Metal verify returned neither logits nor argmax ids");
7153                return None;
7154            }
7155            let ids: Vec<u32> = (0..b)
7156                .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
7157                .collect();
7158            while a < k_spec && ids[a] == drafts[a] {
7159                a += 1;
7160            }
7161            ids
7162        };
7163        spec_stamp("acc");
7164        if spec_dbg {
7165            eprintln!(
7166                "spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
7167                drafts, ids
7168            );
7169        }
7170        // CMF_METAL_VERIFY_CHECK=2: the commit oracle — plain-forward the
7171        // a+1 accepted tokens from a snapshot, then diff the replayed GDN
7172        // states and the appended K/V rows against that.
7173        #[cfg(target_os = "macos")]
7174        let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
7175            && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
7176        {
7177            let snap: Vec<Vec<f32>> = self
7178                .kv_cache
7179                .layers
7180                .iter()
7181                .map(|l| l.linear_state.clone())
7182                .collect();
7183            let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
7184            let toks: Vec<u32> = std::iter::once(t_next)
7185                .chain(drafts.iter().copied())
7186                .collect();
7187            let want_save = self.graph_want_logits;
7188            self.graph_want_logits = false;
7189            for (i, &t) in toks.iter().take(a + 1).enumerate() {
7190                let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
7191                let _ = self.graph_logits.take();
7192            }
7193            self.graph_want_logits = want_save;
7194            let plain_states: Vec<Vec<f32>> = self
7195                .kv_cache
7196                .layers
7197                .iter()
7198                .map(|l| l.linear_state.clone())
7199                .collect();
7200            let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7201            let mut rows = Vec::new();
7202            for (li, (l, n0)) in self
7203                .kv_cache
7204                .layers
7205                .iter_mut()
7206                .zip(attn_lens.iter())
7207                .enumerate()
7208            {
7209                let extra = l.seq_len.saturating_sub(*n0);
7210                if extra > 0 {
7211                    let mut kk = Vec::new();
7212                    let mut vv = Vec::new();
7213                    for g in 0..nkv {
7214                        kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7215                        vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7216                    }
7217                    rows.push((li, kk, vv));
7218                    l.truncate_last(extra);
7219                    crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
7220                }
7221            }
7222            for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7223                if l.linear_state.len() == st.len() {
7224                    l.linear_state.copy_from_slice(&st);
7225                } else {
7226                    l.linear_state = st;
7227                }
7228            }
7229            Some((plain_states, rows))
7230        } else {
7231            None
7232        };
7233        let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
7234        // Metal: the MTP cache cut and the round's warm-up SUBMIT come
7235        // BEFORE the trunk commit, so the warm-up's command buffer is
7236        // queued ahead of the GDN replay (second queue) and its wait
7237        // below no longer sits behind the replay — measured: the warm-up's
7238        // wait grew with the accepted count exactly like the replay does
7239        // (8 ms at a=1, 17 ms at a=3, 25 ms at a=5 for ~2 ms of its own
7240        // work). The replay now overlaps the warm-up's readback, the
7241        // round's return and the next draft chain.
7242        #[cfg(target_os = "macos")]
7243        let mut warm_pending: Option<MetalWarmPending> = None;
7244        #[cfg(target_os = "macos")]
7245        if metal_native {
7246            m.kv.truncate_last(k_spec.saturating_sub(1));
7247            if self.mtp_graph_mode == Some(true) {
7248                // the mirror rows below the cut are the CPU rows: re-point,
7249                // no re-upload
7250                crate::gpu_metal::kv_mirror_set_stored(
7251                    self.mtp_kv_id(),
7252                    Self::MTP_LAYER_BASE,
7253                    m.kv.seq_len,
7254                );
7255                if !warm_off && a > 0 {
7256                    let pairs: Vec<(&[f32], u32)> = (0..a)
7257                        .map(|j| {
7258                            (
7259                                &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
7260                                ids[j],
7261                            )
7262                        })
7263                        .collect();
7264                    warm_pending = self.mtp_warm_batch_submit(m, &pairs, next_pos);
7265                }
7266            }
7267            spec_stamp("c.wsub");
7268        }
7269        // a fully-accepted round needs no restore: every input was real.
7270        #[cfg(target_os = "macos")]
7271        if metal_native {
7272            // the Metal verify never wrote its states: the commit replays the
7273            // accepted prefix into the CPU owners and appends the K/V rows
7274            if !self.metal_verify_commit(a) {
7275                self.clear_sequence_state();
7276                self.graph_failed
7277                    .store(true, std::sync::atomic::Ordering::Relaxed);
7278                self.cancel
7279                    .store(true, std::sync::atomic::Ordering::Relaxed);
7280                tracing::error!("Metal verify state/KV handoff failed after admission");
7281                return None;
7282            }
7283            if let Some((plain_states, rows)) = commit_ref {
7284                crate::gpu_metal::queue_fence();
7285                // the commit's replay runs on the second queue: collect it
7286                // before the oracle reads the CPU owners it writes into
7287                let _ = crate::gpu_metal::wait_replay();
7288                let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7289                let mut worst_s = 0f32;
7290                let mut worst_li = 0usize;
7291                for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
7292                    if l.linear_state.len() != ps.len() || ps.is_empty() {
7293                        continue;
7294                    }
7295                    let d = l
7296                        .linear_state
7297                        .iter()
7298                        .zip(ps)
7299                        .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7300                    let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
7301                    let rel = d / n.max(1e-6);
7302                    if rel > worst_s {
7303                        worst_s = rel;
7304                        worst_li = li;
7305                    }
7306                }
7307                let mut worst_k = 0f32;
7308                for (li, kk, vv) in &rows {
7309                    let l = &self.kv_cache.layers[*li];
7310                    let n0 = l.seq_len - (kk.len() / (nkv * hd));
7311                    let mut ck = Vec::new();
7312                    let mut cv = Vec::new();
7313                    for g in 0..nkv {
7314                        ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7315                        cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7316                    }
7317                    if ck.len() == kk.len() {
7318                        let dk = ck
7319                            .iter()
7320                            .zip(kk)
7321                            .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7322                        let dv = cv
7323                            .iter()
7324                            .zip(vv)
7325                            .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7326                        worst_k = worst_k.max(dk).max(dv);
7327                    } else {
7328                        eprintln!(
7329                            "commit-check L{li}: kv row count mismatch {} vs {}",
7330                            ck.len(),
7331                            kk.len()
7332                        );
7333                    }
7334                }
7335                eprintln!(
7336                    "commit-check a={a}: worst GDN state rel-max diff {worst_s:.2e} (L{worst_li}) | worst K/V row abs diff {worst_k:.4}"
7337                );
7338            }
7339        }
7340        if !metal_native && a + 1 < b {
7341            let expected_gdn_layers = self.graph_gdn_layer_count();
7342            if expected_gdn_layers > 0
7343                && !crate::gpu::gdn_spec_restore(self.graph_kv_id, a, next_pos, expected_gdn_layers)
7344            {
7345                self.clear_sequence_state();
7346                self.graph_failed
7347                    .store(true, std::sync::atomic::Ordering::Relaxed);
7348                self.cancel
7349                    .store(true, std::sync::atomic::Ordering::Relaxed);
7350                tracing::error!("GDN speculative restore failed after verify");
7351                return None;
7352            }
7353        }
7354        if !metal_native && !self.rewind_trunk_graph_mirrors(next_pos + a + 1) {
7355            // The verify graph committed the full batch, but one of its
7356            // persistent Full-attention mirrors could not be re-pointed to
7357            // the accepted prefix.  Treat that as terminal state failure;
7358            // an exact CPU fallback would otherwise consume stale GDN/KV.
7359            self.clear_sequence_state();
7360            self.graph_failed
7361                .store(true, std::sync::atomic::Ordering::Relaxed);
7362            self.cancel
7363                .store(true, std::sync::atomic::Ordering::Relaxed);
7364            tracing::error!("trunk graph KV rewind failed after speculative verify");
7365            return None;
7366        }
7367        *accepted += a;
7368        // MTP cache: keep the first draft row (its inputs were real), drop
7369        // the chain's, then append the verified pairs the round produced.
7370        // Each of those is a whole MTP block on the per-op path and they
7371        // cost 5.8 ms of a 69 ms round at k=3 — a third of what the
7372        // round's own draft costs. PRICED, and they earn it: skipping
7373        // them (`CMF_SPEC_WARM=0`) drops acceptance from 89% to 81% at
7374        // k=3 and 85% to 74% at k=4, and the tok/s goes nowhere at k=3
7375        // (50.3 against 50.5) and backwards at k=4 (48.1 against 50.1).
7376        // The knob stays so the next person can re-price it after the
7377        // warms are batched instead of assuming either way.
7378        if !metal_native {
7379            // (Metal cut its MTP cache before the trunk commit, above)
7380            m.kv.truncate_last(k_spec.saturating_sub(1));
7381        }
7382        spec_stamp("c.trunc");
7383        if !metal_native
7384            && self.mtp_graph_mode == Some(true)
7385            && !self.rewind_mtp_graph_mirror(next_pos)
7386        {
7387            // The graph draft was admitted, so inability to move its cursor
7388            // back to the real anchor is a state failure, not a capability
7389            // refusal.  Do not warm or continue with a stale mirror.
7390            self.clear_sequence_state();
7391            self.graph_failed
7392                .store(true, std::sync::atomic::Ordering::Relaxed);
7393            self.cancel
7394                .store(true, std::sync::atomic::Ordering::Relaxed);
7395            tracing::error!("MTP graph mirror rewind failed after verify commit");
7396            return None;
7397        }
7398        if !warm_off && a > 0 {
7399            // Graph arm: all accepted pairs in ONE batched run over the
7400            // MTP block; the token graph one by one if the batch declines.
7401            let mut warmed = false;
7402            #[cfg(target_os = "macos")]
7403            if metal_native && self.mtp_graph_mode == Some(true) {
7404                // the batched warm-up was submitted before the trunk
7405                // commit: collect it here; one by one on the token graph
7406                // if it declined (or failed)
7407                warmed = match warm_pending.take() {
7408                    Some(p) => self.mtp_warm_batch_finish(m, p),
7409                    None => false,
7410                };
7411                if !warmed {
7412                    warmed = true;
7413                    for j in 0..a {
7414                        let row =
7415                            hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
7416                        if self
7417                            .mtp_step_metal(m, &row, ids[j], next_pos + j, false)
7418                            .is_none()
7419                        {
7420                            warmed = false;
7421                            break;
7422                        }
7423                    }
7424                }
7425            }
7426            if !warmed && self.mtp_graph_mode != Some(false) && !metal_native {
7427                let rows: Vec<Vec<f32>> = (0..a)
7428                    .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
7429                    .collect();
7430                let pairs: Vec<(&[f32], u32)> = rows
7431                    .iter()
7432                    .zip(ids.iter())
7433                    .map(|(r, &t)| (r.as_slice(), t))
7434                    .collect();
7435                match self.mtp_warm_prefill_pairs(m, &pairs, next_pos) {
7436                    Ok(()) => warmed = true,
7437                    Err(err) => {
7438                        // A warm-up failure after graph admission cannot
7439                        // fall back to `mtp_warm`: the detached CPU cache is
7440                        // not authoritative for the device mirror.  Mark it
7441                        // terminal so the generation caller clears state and
7442                        // returns instead of drafting from stale attention.
7443                        tracing::error!("{err}");
7444                        self.clear_sequence_state();
7445                        self.graph_failed
7446                            .store(true, std::sync::atomic::Ordering::Relaxed);
7447                        self.cancel
7448                            .store(true, std::sync::atomic::Ordering::Relaxed);
7449                        return None;
7450                    }
7451                }
7452            }
7453            if !warmed {
7454                for j in 0..a {
7455                    let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
7456                    let row = row.to_vec();
7457                    self.mtp_warm(m, &row, ids[j], next_pos + j);
7458                }
7459            }
7460        }
7461        // The sampler's contract: logits of the LAST verified position —
7462        // unless a rejected draft already drew the correction, in which
7463        // case the loop top commits that token and samples nothing.
7464        spec_stamp("c.warm");
7465        if let Some(c) = forced {
7466            self.spec_forced = Some(c);
7467            self.graph_logits = None;
7468        } else if greedy_dev && logits.is_empty() {
7469            // the row's argmax IS the token the loop top would pick from
7470            // it (plain greedy, no penalties): commit it as forced
7471            self.spec_forced = Some(ids[a]);
7472            self.graph_logits = None;
7473        } else {
7474            let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
7475            row.resize(self.vocab_size, 0.0);
7476            if let Some(c) = self.final_softcap {
7477                for l in row.iter_mut() {
7478                    *l = c * (*l / c).tanh();
7479                }
7480            }
7481            self.graph_logits = Some(row);
7482        }
7483        let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
7484        spec_stamp("c.row");
7485        // Three phases, not two. The round's wall clock was 4 ms longer
7486        // than draft+verify and the difference had nowhere to be seen:
7487        // the accepted prefix re-runs the MTP block once per token to
7488        // keep the draft head's attention cache warm, and the GDN state
7489        // rolls back on any rejection. Both live here, after the verify.
7490        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7491            let end = subs();
7492            eprintln!(
7493                "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
7494                 commit {:.1} ms/{} sub (accepted {a} of {k_spec}, full-head streak {})",
7495                t_draft.as_secs_f64() * 1e3,
7496                sub_draft - sub0,
7497                (t_verify - t_draft).as_secs_f64() * 1e3,
7498                sub_verify - sub_draft,
7499                (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
7500                end - sub_verify,
7501                self.draft_full_streak,
7502            );
7503        }
7504        // Native Metal's verify tile is flat in b (eight rows for the price
7505        // of one), so a shorter round only forfeits tokens — measured on
7506        // the M4: an essay round at k=2 still verified in 260 ms. The
7507        // adaptation is for cards whose verify grows with the rows.
7508        if k_env.is_none() && !metal_native && !k_capped {
7509            // Slow average and a wide band: a fast one oscillated 2↔3 on
7510            // an essay every other round (measured), which forfeits the
7511            // draft it just paid for.
7512            let f = a as f32 / k_spec.max(1) as f32;
7513            self.spec_acc_ewma += 0.2 * (f - self.spec_acc_ewma);
7514            let mut k_next = k_spec;
7515            if self.spec_acc_ewma >= 0.75 && k_spec < k_max {
7516                k_next = k_spec + 1;
7517            } else if self.spec_acc_ewma < 0.4 && k_spec > 2 {
7518                k_next = k_spec - 1;
7519            }
7520            if k_next != k_spec {
7521                self.spec_acc_ewma = 0.6;
7522                if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7523                    eprintln!("spec-k: {k_spec} → {k_next}");
7524                }
7525            }
7526            self.spec_k_adapt = Some(k_next);
7527        }
7528        spec_stamp("end");
7529        Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
7530    }
7531
7532    /// Micro-benchmark: two single-position forwards vs one fused pair
7533    /// from the current cache state (KV rewound after each probe).
7534    /// Returns (two_singles_ms, fused_pair_ms) per probe, or the (0, 0)
7535    /// sentinel when this model has no pair path to measure — the same
7536    /// answer the o1 arm gives, and the bench prints it the same way.
7537    /// (An architecture that loads its own layers leaves `weights.layers`
7538    /// empty; walking it here was an index panic, found by `bench` on
7539    /// deepseek_v4.)
7540    pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
7541        if !self.pair_supported() {
7542            return (0.0, 0.0);
7543        }
7544        // This is a host-side pair micro-benchmark. It truncates the host KV
7545        // after every probe, so letting the whole-token graph participate
7546        // would leave its device GDN/KV mirror ahead of the next probe and
7547        // poison the process-wide graph verdict before the real generation
7548        // benchmark starts. Keep the existing per-op/GPU arithmetic while
7549        // suppressing only the stateful token graph for this measurement.
7550        let graph_env = std::env::var_os("CMF_GPU_WGPU_GRAPH");
7551        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
7552        let emb1 = self.embed_single(1);
7553        let emb2 = self.embed_single(2);
7554        let pos = self.kv_cache.seq_len();
7555
7556        let t0 = std::time::Instant::now();
7557        for _ in 0..iters {
7558            let _ = self.forward_layers(&emb1, pos, None);
7559            let _ = self.forward_layers(&emb2, pos + 1, None);
7560            for l in &mut self.kv_cache.layers {
7561                l.truncate_last(2);
7562            }
7563        }
7564        let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7565
7566        let t1 = std::time::Instant::now();
7567        for _ in 0..iters {
7568            let _ = self.forward_pair(&emb1, &emb2, pos);
7569            for l in &mut self.kv_cache.layers {
7570                l.truncate_last(2);
7571            }
7572        }
7573        let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7574        match graph_env {
7575            Some(value) => unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", value) },
7576            None => unsafe { std::env::remove_var("CMF_GPU_WGPU_GRAPH") },
7577        }
7578        (singles_ms, pair_ms)
7579    }
7580
7581    /// Fused two-position forward: weight rows are streamed from memory
7582    /// once per layer for both positions. Full layers → fused GQA pair;
7583    /// linear layers → vmf_phase pair (lane 2 state is tentative in the
7584    /// per-layer scratch until the draft is accepted).
7585    /// Whether the fused two-position path covers every layer kind in
7586    /// this model. MLA and KDA run per position (their pair arms are
7587    /// unreachable); the seq prefill falls back to singles for them.
7588    fn pair_supported(&self) -> bool {
7589        // An EMPTY layer stack means the architecture loaded its own and
7590        // this path has nothing to walk. Checking that directly, rather
7591        // than naming each such architecture, is what makes the guard hold
7592        // for the next one: `any()` over no layers is false, so a
7593        // feature-by-feature test says "supported" for a model that has no
7594        // layers here at all.
7595        !self.weights.layers.is_empty()
7596            && self.g3n.is_none()
7597            && !self
7598                .weights
7599                .layers
7600                .iter()
7601                .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
7602    }
7603
7604    fn forward_pair(
7605        &mut self,
7606        emb1: &[f32],
7607        emb2: &[f32],
7608        position: usize,
7609    ) -> (Vec<f32>, Vec<f32>) {
7610        // A two-token prompt starts here, not in the layer walk: decide the
7611        // MiMo placement before the pair's per-op MoE uploads any expert.
7612        self.mimo_moe_prepare();
7613        let mut h1 = emb1.to_vec();
7614        let mut h2 = emb2.to_vec();
7615        let (_nkv, _hd, hs, _rd, eps) = (
7616            self.num_kv_heads,
7617            self.head_dim,
7618            self.hidden_size,
7619            self.rotary_dim,
7620            self.rms_eps,
7621        );
7622        let pool = self.pool.clone();
7623
7624        for li in 0..self.num_layers {
7625            let lw = &self.weights.layers[self.phys_layer(li)];
7626            // Norms into pipeline scratch (4 allocs/layer on the MTP
7627            // decode hot path before this).
7628            inference::rms_norm_into(
7629                &h1,
7630                &lw.input_norm,
7631                self.rms_eps,
7632                self.norm_style,
7633                &mut self.ws.n1,
7634            );
7635            inference::rms_norm_into(
7636                &h2,
7637                &lw.input_norm,
7638                self.rms_eps,
7639                self.norm_style,
7640                &mut self.ws.n2,
7641            );
7642
7643            let (a1, a2) = match &lw.attn {
7644                AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
7645                AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
7646                AttnKind::Bounded(w) => {
7647                    // Two sequential positions of the bounded operator
7648                    // (the ring's causal order is the pair's order).
7649                    let rope = self
7650                        .bounded_rope
7651                        .clone()
7652                        .expect("bounded layer without an installed rotation table");
7653                    let cfg = crate::bounded::BoundedAttnCfg {
7654                        num_heads: self.num_heads,
7655                        num_kv_heads: self.num_kv_heads,
7656                        head_dim: self.head_dim,
7657                        hidden_size: hs,
7658                        scale: self.attn_scale,
7659                        rope: &rope,
7660                        pool: pool.as_deref(),
7661                    };
7662                    let a1 = crate::bounded::bounded_attention(
7663                        &self.ws.n1,
7664                        w,
7665                        &mut self.kv_cache.layers[li],
7666                        &cfg,
7667                    );
7668                    let a2 = crate::bounded::bounded_attention(
7669                        &self.ws.n2,
7670                        w,
7671                        &mut self.kv_cache.layers[li],
7672                        &cfg,
7673                    );
7674                    (a1, a2)
7675                }
7676                AttnKind::Linear(w) => {
7677                    let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
7678                    let layer = &mut self.kv_cache.layers[li];
7679                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7680                    vmf_phase_pair(
7681                        &self.ws.n1,
7682                        &self.ws.n2,
7683                        w,
7684                        &cfg,
7685                        state,
7686                        scratch,
7687                        self.pool.as_deref(),
7688                    )
7689                }
7690                AttnKind::LinearGdn(w) => {
7691                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
7692                    let layer = &mut self.kv_cache.layers[li];
7693                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7694                    gdn_pair(
7695                        &self.ws.n1,
7696                        &self.ws.n2,
7697                        w,
7698                        &cfg,
7699                        state,
7700                        scratch,
7701                        self.pool.as_deref(),
7702                    )
7703                }
7704                AttnKind::ShortConv(w) => {
7705                    let cfg = self
7706                        .short_conv_cfg
7707                        .expect("short-conv layer without short_conv_cfg");
7708                    let layer = &mut self.kv_cache.layers[li];
7709                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7710                    short_conv_pair(
7711                        &self.ws.n1,
7712                        &self.ws.n2,
7713                        w,
7714                        &cfg,
7715                        state,
7716                        scratch,
7717                        self.pool.as_deref(),
7718                    )
7719                }
7720                AttnKind::Full {
7721                    wq,
7722                    wk,
7723                    wv,
7724                    wo,
7725                    q_norm,
7726                    k_norm,
7727                    output_gate,
7728                    softplus_gate,
7729                    bias,
7730                } => {
7731                    let inv_freq_l = self.layer_inv_freq(li);
7732                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
7733                    let cfg = QwenAttnCfg {
7734                        num_heads: self.layer_num_heads(li),
7735                        num_kv_heads: nkv_l,
7736                        head_dim: hd_l,
7737                        hidden_size: hs,
7738                        position,
7739                        inv_freq: &inv_freq_l,
7740                        rotary_dim: rd_l,
7741                        scale: self.attn_scale,
7742                        softcap: self.attn_softcap,
7743                        window: self.layer_window(li),
7744                        v_norm: self.attn_v_norm,
7745                        qk_norm_after_rope: self.qk_norm_after_rope,
7746                        gate_sigmoid: self.proj_gate_sigmoid,
7747                        q_norm: q_norm.as_deref(),
7748                        k_norm: k_norm.as_deref(),
7749                        output_gate: *output_gate,
7750                        softplus_gate: softplus_gate
7751                            .as_ref()
7752                            .map(|(gate, per_head)| (gate, *per_head)),
7753                        rope_scale: self.layer_rope_scale(li),
7754                        bias: bias
7755                            .as_ref()
7756                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7757                        rms_eps: eps,
7758                        norm_style: self.norm_style,
7759                        pool: pool.as_deref(),
7760                        v_head_dim: self.layer_v_dim(li),
7761                    };
7762                    attention::qwen_attention_pair(
7763                        &self.ws.n1,
7764                        &self.ws.n2,
7765                        wq,
7766                        wk,
7767                        wv,
7768                        wo,
7769                        &mut self.kv_cache.layers[li],
7770                        &cfg,
7771                    )
7772                }
7773            };
7774            let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
7775                Some(w) => (
7776                    inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
7777                    inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
7778                ),
7779                None => (a1, a2),
7780            };
7781            for i in 0..self.hidden_size {
7782                h1[i] += a1[i];
7783                h2[i] += a2[i];
7784            }
7785            let (mut a1, mut a2) = (a1, a2);
7786            attention::recycle_buf(&mut a1);
7787            attention::recycle_buf(&mut a2);
7788
7789            let lw = &self.weights.layers[self.phys_layer(li)];
7790            inference::rms_norm_into(
7791                &h1,
7792                &lw.post_norm,
7793                self.rms_eps,
7794                self.norm_style,
7795                &mut self.ws.p1,
7796            );
7797            inference::rms_norm_into(
7798                &h2,
7799                &lw.post_norm,
7800                self.rms_eps,
7801                self.norm_style,
7802                &mut self.ws.p2,
7803            );
7804            let (f1, f2) = match &lw.ffn {
7805                // Dual-branch layers need the raw residuals — run the
7806                // two positions through the same fn decode uses.
7807                FfnKind::DenseMoe(dm) => (
7808                    dense_moe_ffn(
7809                        dm,
7810                        &self.ws.p1,
7811                        &h1,
7812                        self.rms_eps,
7813                        self.norm_style,
7814                        self.pool.as_deref(),
7815                    ),
7816                    dense_moe_ffn(
7817                        dm,
7818                        &self.ws.p2,
7819                        &h2,
7820                        self.rms_eps,
7821                        self.norm_style,
7822                        self.pool.as_deref(),
7823                    ),
7824                ),
7825                FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => (
7826                    moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p1, self.pool.as_deref()),
7827                    moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p2, self.pool.as_deref()),
7828                ),
7829                _ => ffn_forward_pair(
7830                    &lw.ffn,
7831                    &self.ws.p1,
7832                    &self.ws.p2,
7833                    self.pool.as_deref(),
7834                    None,
7835                ),
7836            };
7837            let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
7838                Some(w) => (
7839                    inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
7840                    inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
7841                ),
7842                None => (f1, f2),
7843            };
7844            for i in 0..self.hidden_size {
7845                h1[i] += f1[i];
7846                h2[i] += f2[i];
7847            }
7848            let (mut f1, mut f2) = (f1, f2);
7849            attention::recycle_buf(&mut f1);
7850            attention::recycle_buf(&mut f2);
7851            if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
7852                for i in 0..self.hidden_size {
7853                    h1[i] *= sc;
7854                    h2[i] *= sc;
7855                }
7856            }
7857            // Looped Transformer: apply final norm at the end of each loop iteration.
7858            if self.is_loop_end(li) && li + 1 < self.num_layers {
7859                h1 = inference::rms_norm(
7860                    &h1,
7861                    &self.weights.final_norm,
7862                    self.rms_eps,
7863                    self.norm_style,
7864                );
7865                h2 = inference::rms_norm(
7866                    &h2,
7867                    &self.weights.final_norm,
7868                    self.rms_eps,
7869                    self.norm_style,
7870                );
7871            }
7872        }
7873        // Real O(1) prefill pairs may also carry tentative lane-2 recurrent
7874        // state. Commit it before publishing the transition epoch so the
7875        // next serial/device row cannot observe a new attention epoch with an
7876        // old GDN state. Speculative pairs run only when O(1) is inactive and
7877        // retain their existing caller-controlled commit/rollback semantics.
7878        if self.o1_active() {
7879            self.commit_linear_scratch();
7880        }
7881        self.o1_progress();
7882        (h1, h2)
7883    }
7884
7885    /// Commit lane-2 linear states after an accepted draft.
7886    fn commit_linear_scratch(&mut self) {
7887        for layer in &mut self.kv_cache.layers {
7888            if !layer.linear_scratch.is_empty() {
7889                std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
7890                layer.linear_scratch.clear();
7891            }
7892        }
7893    }
7894
7895    /// Forward a full id sequence from a fresh cache and return the
7896    /// logits after the last position (golden-parity harness, bench).
7897    pub fn forward_ids(
7898        &mut self,
7899        ids: &[u32],
7900        task_mask: Option<&TaskMask>,
7901    ) -> Result<Vec<f32>, String> {
7902        #[cfg(target_os = "macos")]
7903        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7904        if ids.is_empty() {
7905            return Err("empty id sequence".to_string());
7906        }
7907        self.clear_sequence_state();
7908        self.check_forward_graph("forward_ids setup", 0)?;
7909        if task_mask.is_none() {
7910            self.o1_begin();
7911        }
7912        let mut hidden = vec![0.0f32; self.hidden_size];
7913        let mut pos = 0usize;
7914        if let Some(b) = &mut self.dsv41 {
7915            let pool = self.pool.clone();
7916            let mut logits = Vec::new();
7917            crate::dsv41::forward_chunk(
7918                &b.0,
7919                &b.1,
7920                &b.2,
7921                &mut b.3,
7922                ids,
7923                0,
7924                pool.as_deref(),
7925                &mut logits,
7926            );
7927            if let Err(err) = self.o1_seal_checked() {
7928                self.clear_sequence_state();
7929                return Err(err);
7930            }
7931            return Ok(logits);
7932        }
7933        // Same routing predicate generation uses. Two reasons it must be
7934        // the same one: (1) a GDN hybrid's recurrent state is GPU-
7935        // resident, and a batched CPU prefill would build it on the host
7936        // only — decode then reads buffers the prefill never wrote;
7937        // (2) bench times THIS function and calls the result "prefill",
7938        // so a different path here reports a number production never
7939        // sees (W2 on 2×5090: 8.7 tok/s reported against 125 real).
7940        if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
7941            // prefill-GEMM in chunks; only the last position's hidden is
7942            // needed. (o1-compatible: the batch path attends per position
7943            // through qwen_attention, which carries the collection hook.)
7944            let chunk = self.prefill_chunk();
7945            let hs = self.hidden_size;
7946            while pos < ids.len() {
7947                let end = (pos + chunk).min(ids.len());
7948                let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
7949                self.check_forward_graph("forward_ids batched prefill", end - 1)?;
7950                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
7951                pos = end;
7952            }
7953        }
7954        // Same guards as generation's prefill — INCLUDING the graph one.
7955        // The CPU pair walk was intercepting positions that the resident
7956        // token graph would have run itself: on a GDN hybrid over wgpu
7957        // that is 89 ms of host forward against 7 ms of device submit,
7958        // and it made prefill look 12× slower than it is (W2 on an RTX
7959        // 5090, ctx 512: 11.2 tok/s with the walk, 136.6 without).
7960        // CMF_PAIR=0 opts out; a model whose layers live outside
7961        // `weights.layers` has no pair walk to take.
7962        if task_mask.is_none()
7963            && !self.graph_prefill_preferred()
7964            && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
7965            && self.pair_supported()
7966        {
7967            while pos + 1 < ids.len() {
7968                let e1 = self.embed_single(ids[pos]);
7969                let e2 = self.embed_single(ids[pos + 1]);
7970                let (_, h2) = self.forward_pair(&e1, &e2, pos);
7971                self.check_forward_graph("forward_ids pair", pos + 1)?;
7972                self.commit_linear_scratch();
7973                hidden = h2;
7974                pos += 2;
7975            }
7976        }
7977        // Resident Embryo graph: the prompt in chunks of one submit each
7978        // (the same device state and logits as the per-position walk).
7979        if task_mask.is_none() && pos == 0 && ids.len() > 1 {
7980            if let Some(lg) = self.embryo_prefill_chunked(ids, 0) {
7981                self.graph_logits = Some(lg);
7982                hidden = vec![0.0; self.hidden_size];
7983                pos = ids.len();
7984            }
7985        }
7986        while pos < ids.len() {
7987            hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
7988            self.check_forward_graph("forward_ids", pos)?;
7989            pos += 1;
7990        }
7991        if let Some(logits) = self.graph_logits.take() {
7992            // Resident stacks already applied final norm and their head in
7993            // the same submit; do not run a second norm/head over the zero
7994            // hidden sentinel returned by forward_layers_span.
7995            if let Err(err) = self.o1_seal_checked() {
7996                self.clear_sequence_state();
7997                return Err(err);
7998            }
7999            return Ok(logits);
8000        }
8001        // Harness contract: after forward_ids the cache is decode-ready —
8002        // under o1 that means sealed (bench measures the seal as part of
8003        // prefill, honestly).
8004        if let Err(err) = self.o1_seal_checked() {
8005            self.clear_sequence_state();
8006            return Err(err);
8007        }
8008        let normed = inference::rms_norm(
8009            &hidden,
8010            &self.weights.final_norm,
8011            self.rms_eps,
8012            self.norm_style,
8013        );
8014        Ok(self.lm_head_forward(&normed))
8015    }
8016
8017    /// Run the V4.1 stack one token at a time and retain logits for every
8018    /// position. This is a diagnostic surface for comparing a converted
8019    /// checkpoint with a tokenwise reference implementation.
8020    #[doc(hidden)]
8021    pub fn dsv41_serial_logits(&mut self, ids: &[u32]) -> Result<Vec<Vec<f32>>, String> {
8022        #[cfg(target_os = "macos")]
8023        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8024        if ids.is_empty() {
8025            return Err("empty id sequence".to_string());
8026        }
8027        self.clear_sequence_state();
8028        self.dsv41
8029            .as_ref()
8030            .ok_or_else(|| "dsv41 serial logits require a DeepSeek-V4.1 model".to_string())?;
8031        self.o1_begin();
8032        let rows = {
8033            let pool = self.pool.clone();
8034            let b = self
8035                .dsv41
8036                .as_mut()
8037                .expect("dsv41 checked above; state cannot change during forward");
8038            let mut rows = Vec::with_capacity(ids.len());
8039            for (position, &id) in ids.iter().enumerate() {
8040                let mut logits = Vec::new();
8041                crate::dsv41::forward_token(
8042                    &b.0,
8043                    &b.1,
8044                    &b.2,
8045                    &mut b.3,
8046                    id,
8047                    position,
8048                    pool.as_deref(),
8049                    &mut logits,
8050                );
8051                rows.push(logits);
8052            }
8053            rows
8054        };
8055        self.o1_seal();
8056        Ok(rows)
8057    }
8058
8059    /// Teacher-forced perplexity over a token sequence (phase-C gate:
8060    /// honest quant comparisons instead of prompt vibes).
8061    ///
8062    /// Attention is EXACT even on a model whose layers are flagged for
8063    /// the O(1) kernel — scoring the backbone is the default on purpose
8064    /// (it is the yardstick). `nll_ids_o1` scores the CONVERTED model.
8065    pub fn ppl_ids(&mut self, ids: &[u32]) -> Result<f64, String> {
8066        let (nll, cnt) = self.nll_ids_from(ids, 0)?;
8067        Ok((nll / cnt.max(1) as f64).exp())
8068    }
8069
8070    /// DTG-MA calibration pass (Patent 2): run `ids` through the model
8071    /// (CPU path, per position) and return each layer's per-neuron
8072    /// activation mass Σ|silu(gate)·up| — the statistic the task-guided
8073    /// FFN mask is derived from.
8074    pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
8075        self.clear_sequence_state();
8076        FFN_PROBE.with(|p| {
8077            *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
8078        });
8079        crate::gpu::cpu_scope(|| {
8080            for (pos, &id) in ids.iter().enumerate() {
8081                let emb = self.embed_single(id);
8082                let _ = self.forward_layers(&emb, pos, None);
8083            }
8084        });
8085        self.clear_sequence_state();
8086        FFN_PROBE
8087            .with(|p| p.borrow_mut().take())
8088            .unwrap_or_default()
8089    }
8090
8091    /// `probe_ffn_mass` over the BATCHED prefill: same accumulator, one
8092    /// sweep instead of one forward per token. What makes the statistic
8093    /// affordable on a 27B.
8094    pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Result<Vec<Vec<f64>>, String> {
8095        if let Err(err) = self.nll_begin() {
8096            // A recorder can be left by a caller that was interrupted before
8097            // this request entered its scoring block.  Consume it even when
8098            // the preflight failure prevents initialization of a new one.
8099            let _ = FFN_PROBE.with(|p| p.borrow_mut().take());
8100            self.nll_end();
8101            return Err(err);
8102        }
8103        FFN_PROBE.with(|p| {
8104            *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
8105        });
8106        let result: Result<(), String> = (|| {
8107            for chunk in ids.chunks(256) {
8108                if chunk.len() < 2 {
8109                    continue;
8110                }
8111                self.nll_ids_masked(chunk, 0, None)?;
8112            }
8113            Ok(())
8114        })();
8115        self.nll_end();
8116        let probe = FFN_PROBE
8117            .with(|p| p.borrow_mut().take())
8118            .unwrap_or_default();
8119        match result {
8120            Ok(()) => Ok(probe),
8121            Err(err) => {
8122                drop(probe);
8123                Err(err)
8124            }
8125        }
8126    }
8127
8128    /// Teacher-forced PPL with a task mask active (sparse execution) —
8129    /// the quality gate for a DTG-MA-masked skill. Sequential per
8130    /// position: the batched prefill path is dense-only.
8131    pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> Result<f64, String> {
8132        self.nll_begin()?;
8133        let result: Result<f64, String> = (|| {
8134            let mut nll = 0f64;
8135            let mut cnt = 0usize;
8136            let mut hidden = vec![0f32; self.hidden_size];
8137            for (pos, &id) in ids.iter().enumerate() {
8138                if pos > 0 {
8139                    inference::rms_norm_into(
8140                        &hidden,
8141                        &self.weights.final_norm,
8142                        self.rms_eps,
8143                        self.norm_style,
8144                        &mut self.ws.n1,
8145                    );
8146                    let mut logits = self.lm_head_forward(&self.ws.n1);
8147                    let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
8148                    let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
8149                    let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
8150                    nll -= p.max(1e-300).ln();
8151                    cnt += 1;
8152                    attention::recycle_buf(&mut logits);
8153                }
8154                let emb = self.embed_single(id);
8155                hidden = self.forward_layers(&emb, pos, Some(mask));
8156                self.nll_check_graph("masked serial forward", pos)?;
8157                // Consume a possible graph logits side channel before the
8158                // next row.  Masked scoring normally disables that route,
8159                // but stale channel state must never survive a request.
8160                let _ = self.graph_logits.take();
8161            }
8162            Ok((nll / cnt.max(1) as f64).exp())
8163        })();
8164        self.nll_end();
8165        result
8166    }
8167
8168    /// Teacher-forced NLL sum + scored-token count over positions
8169    /// `start..len-1`, attention EXACT. Positions below `start` still
8170    /// run — they are the context — they are just not scored, so this
8171    /// pairs with `nll_ids_o1(ids, start)` over the very same tokens.
8172    ///
8173    /// Returning (nll, cnt) rather than a ppl is what lets a windowed
8174    /// caller combine windows before the exp, so every scored token
8175    /// weighs the same regardless of how the windows are cut.
8176    /// `nll_ids_from` with a task mask held active at every position.
8177    ///
8178    /// The batched prefill path does not thread masks, so this walks the
8179    /// per-position forward — slower, but it scores the file exactly the
8180    /// way `run --task` will serve it, which is the point of the gate
8181    /// that calls it. With `None` it defers to the fast path.
8182    /// Masked scoring rides the SAME batched sweep as unmasked scoring —
8183    /// the masked-inference fast path: `prefill_batch_masked` lands the
8184    /// per-visit FFN rows on the activations inside the fused arms. The
8185    /// per-position loop below remains only as the no-batch fallback.
8186    pub fn nll_ids_masked(
8187        &mut self,
8188        ids: &[u32],
8189        start: usize,
8190        task_mask: Option<&TaskMask>,
8191    ) -> Result<(f64, usize), String> {
8192        let task_mask = self.drop_open_mask(task_mask);
8193        self.nll_ids_inner(ids, start, task_mask)
8194    }
8195
8196    pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> Result<(f64, usize), String> {
8197        self.nll_ids_inner(ids, start, None)
8198    }
8199
8200    fn nll_ids_inner(
8201        &mut self,
8202        ids: &[u32],
8203        start: usize,
8204        task_mask: Option<&TaskMask>,
8205    ) -> Result<(f64, usize), String> {
8206        self.nll_begin()?;
8207        let result: Result<(f64, usize), String> = (|| {
8208            let mut nll = 0f64;
8209            let mut cnt = 0usize;
8210            // An unmasked quality run with the resident wgpu graph must score
8211            // the same stateful path used by generation.  The layer-major
8212            // GEMM prefill below is a valid CPU/GEMM oracle, but it seeds
8213            // neither the graph's device GDN state nor its device KV mirrors;
8214            // using it here would silently score a different execution.  Keep
8215            // masked scoring on the exact per-position path as before, and
8216            // let the serial arm below drive the graph-aware scorer.
8217            // Only native Metal has a fused graph lm_head contract.  Vulkan
8218            // and other graph backends may expose hidden state without the
8219            // optional logits side channel; preserve their established CPU
8220            // norm/head fallback instead of turning that valid route into a
8221            // hard missing-logits error.
8222            let (graph_quality, fused_head_quality) = nll_graph_policy(
8223                task_mask.is_none(),
8224                self.graph_prefill_preferred(),
8225                crate::gpu::q1_force(),
8226            );
8227            self.graph_head_required = fused_head_quality;
8228            self.graph_want_logits = fused_head_quality;
8229            #[cfg(target_os = "macos")]
8230            if graph_quality && std::env::var("CMF_METAL_BATCH_NLL").as_deref() != Ok("0") {
8231                match self.nll_batch_metal(ids, start) {
8232                    MetalBatchNllOutcome::Completed(nll, count) => {
8233                        return Ok((nll, count));
8234                    }
8235                    MetalBatchNllOutcome::Declined => {}
8236                    MetalBatchNllOutcome::Failed(err) => return Err(err),
8237                }
8238            }
8239            // CMF_NLL_SERIAL=1 scores position by position through the
8240            // decode path (the wgpu token graph where it admits the model)
8241            // — the perplexity gate of a decode kernel, which the batched
8242            // prefill arm below never runs.
8243            let force_serial = std::env::var("CMF_NLL_SERIAL").as_deref() == Ok("1");
8244            if self.can_prefill_batched() && !graph_quality && !force_serial {
8245                // prefill-GEMM: layer-major position chunks, lm_head batched
8246                // (254MB lm_head read once per chunk, not per position).
8247                // The layer chunk is large (grouping positions by MoE experts
8248                // wins with size), lm_head in sub-blocks (logit buffer
8249                // 32×vocab ≈ 32MB instead of 128×).
8250                const CHUNK: usize = 128;
8251                const LM_SUB: usize = 32;
8252                let n = ids.len().saturating_sub(1);
8253                let hs = self.hidden_size;
8254                let rows = self.weights.lm_head.rows();
8255                let mut pos = 0usize;
8256                let state_trace = std::env::var("CMF_STATE_TRACE").is_ok();
8257                while pos < n {
8258                    let end = (pos + CHUNK).min(n);
8259                    let bsz = end - pos;
8260                    let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
8261                    self.nll_check_graph("batched prefill", pos)?;
8262                    if state_trace && end % 256 == 0 {
8263                        self.trace_recurrent_state(end);
8264                    }
8265                    let mut k0 = 0usize;
8266                    while k0 < bsz {
8267                        let k1 = (k0 + LM_SUB).min(bsz);
8268                        let sb = k1 - k0;
8269                        // Sub-block entirely below the scored range: the KV
8270                        // it just built is all this pass needed from it.
8271                        if pos + k1 <= start {
8272                            k0 = k1;
8273                            continue;
8274                        }
8275                        let mut normed = vec![0.0f32; sb * hs];
8276                        for k in 0..sb {
8277                            let r = inference::rms_norm(
8278                                &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
8279                                &self.weights.final_norm,
8280                                self.rms_eps,
8281                                self.norm_style,
8282                            );
8283                            normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
8284                        }
8285                        let mut logits = vec![0.0f32; sb * rows];
8286                        self.weights
8287                            .lm_head
8288                            .matmat(&normed, sb, &mut logits, self.pool.as_deref());
8289                        for k in 0..sb {
8290                            if pos + k0 + k < start {
8291                                continue;
8292                            }
8293                            self.nll_check_graph("batched score row", pos + k0 + k)?;
8294                            let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
8295                            if let Some(mu) = self.logit_multiplier {
8296                                for v in lg.iter_mut() {
8297                                    *v *= mu;
8298                                }
8299                            }
8300                            // Gemma-class final-logit soft-capping: the
8301                            // decode paths apply it; scoring must too, or
8302                            // the uncapped softmax misprices every token.
8303                            if let Some(c) = self.final_softcap {
8304                                for v in lg.iter_mut() {
8305                                    *v = c * (*v / c).tanh();
8306                                }
8307                            }
8308                            // Cortiq Embryo hierarchical head: same correction
8309                            // the decode path applies (lm_head_forward).
8310                            if let Some(cm) = self.head_clusters.clone() {
8311                                self.hierarchical_head_logprobs(
8312                                    &normed[k * hs..(k + 1) * hs],
8313                                    &cm,
8314                                    lg,
8315                                );
8316                            }
8317                            let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
8318                            let target = ids[pos + k0 + k + 1] as usize;
8319                            let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8320                            let lse: f64 = lg
8321                                .iter()
8322                                .map(|&v| ((v - max) as f64).exp())
8323                                .sum::<f64>()
8324                                .ln()
8325                                + max as f64;
8326                            nll += lse - lg[target] as f64;
8327                            cnt += 1;
8328                            if std::env::var("CMF_PPL_TRACE").is_ok() {
8329                                let top = lg
8330                                    .iter()
8331                                    .enumerate()
8332                                    .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8333                                    .map(|(i, _)| i)
8334                                    .unwrap_or(0);
8335                                eprintln!(
8336                                    "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
8337                                    pos + k0 + k,
8338                                    target,
8339                                    lse - lg[target] as f64,
8340                                    top,
8341                                    lg[target],
8342                                    lg[top]
8343                                );
8344                            }
8345                        }
8346                        k0 = k1;
8347                    }
8348                    pos = end;
8349                }
8350                return Ok((nll, cnt));
8351            }
8352            for pos in 0..ids.len().saturating_sub(1) {
8353                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
8354                self.nll_check_graph("serial forward", pos)?;
8355                // Architectures whose head lives inside their own stack return
8356                // the logits out of band and a zero hidden — DeepSeek-V4 folds
8357                // its hyper-connection copies between the last layer and the
8358                // norm, so it cannot hand back a vector this loop could use.
8359                // Scoring the zeros gave a perplexity of exactly the vocabulary
8360                // size, which is a uniform distribution reported as a
8361                // measurement. `generate` already reads this channel.
8362                let out_of_band = self.graph_logits.take();
8363                if self.graph_head_required && out_of_band.is_none() {
8364                    METAL_GRAPH_HEAD_MISS.fetch_add(
8365                        1,
8366                        std::sync::atomic::Ordering::Relaxed,
8367                    );
8368                    return Err(format!(
8369                        "fused Metal graph head did not complete at NLL position {pos}"
8370                    ));
8371                }
8372                if pos < start {
8373                    continue;
8374                }
8375                let logits = match out_of_band {
8376                    Some(lg) => lg,
8377                    None => {
8378                        let normed = inference::rms_norm(
8379                            &hidden,
8380                            &self.weights.final_norm,
8381                            self.rms_eps,
8382                            self.norm_style,
8383                        );
8384                        // lm_head_forward applies the final-logit softcap itself
8385                        // — capping again here double-squashed gemma-class
8386                        // logits (tanh∘tanh) and reported a flattered ppl.
8387                        self.lm_head_forward(&normed)
8388                    }
8389                };
8390                let target = ids[pos + 1] as usize;
8391                let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8392                let lse: f64 = logits
8393                    .iter()
8394                    .map(|&v| ((v - max) as f64).exp())
8395                    .sum::<f64>()
8396                    .ln()
8397                    + max as f64;
8398                let tok_nll = lse - logits[target] as f64;
8399                if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8400                    let top = logits
8401                        .iter()
8402                        .enumerate()
8403                        .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8404                        .map(|(i, _)| i)
8405                        .unwrap_or(0);
8406                    eprintln!(
8407                        "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8408                        logits[target], logits[top]
8409                    );
8410                }
8411                nll += tok_nll;
8412                cnt += 1;
8413            }
8414            Ok((nll, cnt))
8415        })();
8416        self.nll_end();
8417        result
8418    }
8419
8420    /// Score one post-layer hidden with the same final norm/head path used by
8421    /// decode. Keeping this in one helper is important for the production
8422    /// batch scorer: its rows stop before the final norm, just like the
8423    /// per-position O(1) path below.
8424    fn nll_from_hidden(&mut self, hidden: &[f32], target: u32, pos: usize) -> f64 {
8425        let normed = inference::rms_norm(
8426            hidden,
8427            &self.weights.final_norm,
8428            self.rms_eps,
8429            self.norm_style,
8430        );
8431        // lm_head_forward applies the final-logit softcap itself — capping
8432        // again here double-squashed gemma-class logits in earlier scorers.
8433        let mut logits = self.lm_head_forward(&normed);
8434        let target = target as usize;
8435        let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8436        let lse: f64 = logits
8437            .iter()
8438            .map(|&v| ((v - max) as f64).exp())
8439            .sum::<f64>()
8440            .ln()
8441            + max as f64;
8442        let tok_nll = lse - logits[target] as f64;
8443        if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8444            let top = logits
8445                .iter()
8446                .enumerate()
8447                .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8448                .map(|(i, _)| i)
8449                .unwrap_or(0);
8450            eprintln!(
8451                "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8452                logits[target], logits[top]
8453            );
8454        }
8455        attention::recycle_buf(&mut logits);
8456        tok_nll
8457    }
8458
8459    /// `CMF_STATE_TRACE`: per-layer magnitude of the recurrent record at
8460    /// position `pos` — the whole `linear_state` (vmf: S then the conv
8461    /// ring; GDN: conv ring then S), its recurrent S part alone, the
8462    /// bounded ring, and the last per-position KV row count.  The tool
8463    /// that separated the ~4k perplexity cliff of the 500-step exports
8464    /// (a state that keeps climbing past the trained window) from a
8465    /// runtime boundary; one line per layer, `STATE pos=… layer=…`.
8466    fn trace_recurrent_state(&self, pos: usize) {
8467        let stats = |v: &[f32]| -> (f64, f64) {
8468            if v.is_empty() {
8469                return (0.0, 0.0);
8470            }
8471            let (mut ss, mut mx) = (0f64, 0f64);
8472            for &x in v {
8473                ss += (x as f64) * (x as f64);
8474                mx = mx.max((x as f64).abs());
8475            }
8476            ((ss / v.len() as f64).sqrt(), mx)
8477        };
8478        for (li, l) in self.kv_cache.layers.iter().enumerate() {
8479            let lw = &self.weights.layers[self.phys_layer(li)];
8480            let (kind, s_len) = match &lw.attn {
8481                AttnKind::Linear(_) => (
8482                    "vmf",
8483                    self.vmf_cfg.map(|c| c.state_len()).unwrap_or(0),
8484                ),
8485                AttnKind::LinearGdn(_) => (
8486                    "gdn",
8487                    self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0),
8488                ),
8489                AttnKind::Bounded(_) => ("bounded", 0),
8490                AttnKind::Full { .. } => ("full", 0),
8491                _ => ("other", 0),
8492            };
8493            let (rms, max) = stats(&l.linear_state);
8494            let s_part = if kind == "vmf" {
8495                &l.linear_state[..s_len.min(l.linear_state.len())]
8496            } else if kind == "gdn" {
8497                let ring = l.linear_state.len().saturating_sub(s_len).min(l.linear_state.len());
8498                &l.linear_state[ring..]
8499            } else {
8500                &l.linear_state[..0]
8501            };
8502            let (s_rms, s_max) = stats(s_part);
8503            let (ring_rms, ring_len) = match &l.bounded {
8504                Some(b) => (stats(&b.ring_k).0, b.ring_k.len()),
8505                None => (0.0, 0),
8506            };
8507            eprintln!(
8508                "STATE pos={pos} layer={li} kind={kind} state_len={} rms={rms:.5} max={max:.4} \
8509                 S_rms={s_rms:.5} S_max={s_max:.4} ring_k_rms={ring_rms:.5} ring_len={ring_len} kv_rows={}",
8510                l.linear_state.len(),
8511                l.seq_len
8512            );
8513        }
8514    }
8515
8516    /// Teacher-forced NLL of the CONVERTED model: the O(1) Nyström path
8517    /// is ACTIVE over the scored positions. Returns `Ok((nll sum, scored
8518    /// count))` over `prefill..len-1` and surfaces a post-mutation batch
8519    /// failure instead of returning a partial score.
8520    ///
8521    /// Runtime discipline, deliberately NOT the matrix probe's: the
8522    /// requested prefix plus any required deferred lead-in run the exact
8523    /// prompt pass — that pass is what freezes the landmarks and M — and
8524    /// every post-seal scored position goes through `NystromState::step()`,
8525    /// the same code decode runs.
8526    /// So the landmarks are PREFILL-frozen (what ships), not
8527    /// full-sequence oracles (what the published probe measured). When the
8528    /// requested prefix is shorter than the bounded transition, rows in the
8529    /// exact lead-in are still scored so the shifted target range is stable.
8530    ///
8531    /// Pair with `nll_ids_from(ids, prefill)` for the exact baseline
8532    /// over the identical token set — that ratio is the honest one.
8533    pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> Result<(f64, usize), String> {
8534        // This scorer consumes host hiddens, so never request the optional
8535        // token-graph lm_head side channel. `nll_begin` also consumes a
8536        // prior graph failure and clears only the cancel bit that failure
8537        // raised, leaving a caller-owned cancellation observable.
8538        self.nll_begin()?;
8539        let requested_prefix = (prefill > 0).then_some(prefill);
8540        self.o1_begin_with_prefix(requested_prefix);
8541        let n = ids.len().saturating_sub(1);
8542        let requested_start = prefill.min(n);
8543        // The exact prefix must reach the deferred boundary before a
8544        // collecting layer can convert. Rows between the requested start and
8545        // that boundary remain part of the public NLL range and are scored
8546        // from the same hidden pass below.
8547        let exact_end = if self.o1_active() {
8548            match requested_prefix {
8549                Some(requested) => self.o1_effective_boundary(requested),
8550                None => self
8551                    .o1_cfg
8552                    .as_ref()
8553                    .and_then(|c| crate::nystrom::o1_deferred_boundary(c.w, c.sink)),
8554            }
8555            .unwrap_or(requested_start)
8556            .min(n)
8557        } else {
8558            requested_start
8559        };
8560        let mut nll = 0f64;
8561        let mut cnt = 0usize;
8562
8563        // Exact prompt pass over ids[..exact_end]: the seal consumes its
8564        // q/k/v. Rows at or after requested_start are scored here when the
8565        // bounded lead-in is longer than the caller's requested prefix.
8566        let mut pos = 0usize;
8567        if self.can_prefill_batched() {
8568            const CHUNK: usize = 128;
8569            while pos < exact_end {
8570                let end = (pos + CHUNK).min(exact_end);
8571                let hiddens = self.prefill_batch(&ids[pos..end], pos);
8572                if self
8573                    .graph_failed
8574                    .swap(false, std::sync::atomic::Ordering::Relaxed)
8575                {
8576                    self.cancel
8577                        .store(false, std::sync::atomic::Ordering::Relaxed);
8578                    self.nll_end();
8579                    return Err("GPU graph failed during O(1) NLL prefix".into());
8580                }
8581                for row in 0..end - pos {
8582                    let score_pos = pos + row;
8583                    if score_pos >= requested_start && score_pos < n {
8584                        nll += self.nll_from_hidden(
8585                            &hiddens[row * self.hidden_size..(row + 1) * self.hidden_size],
8586                            ids[score_pos + 1],
8587                            score_pos,
8588                        );
8589                        cnt += 1;
8590                    }
8591                }
8592                pos = end;
8593            }
8594        } else {
8595            while pos < exact_end {
8596                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8597                if self
8598                    .graph_failed
8599                    .swap(false, std::sync::atomic::Ordering::Relaxed)
8600                {
8601                    self.cancel
8602                        .store(false, std::sync::atomic::Ordering::Relaxed);
8603                    self.nll_end();
8604                    return Err("GPU graph failed during O(1) NLL prefix".into());
8605                }
8606                if pos >= requested_start && pos < n {
8607                    nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8608                    cnt += 1;
8609                }
8610                pos += 1;
8611            }
8612        }
8613        self.o1_seal_checked().map_err(|err| {
8614            self.nll_end();
8615            err
8616        })?;
8617
8618        // Reuse the production whole-token batch graph for the post-seal
8619        // suffix when the caller explicitly enabled both routes. This is a
8620        // teacher-forced scorer, so every row is ids[pos] and its target is
8621        // ids[pos + 1]; no speculative tail or rollback state is involved.
8622        // A first Declined is safe to handle with the established serial O(1)
8623        // path. Once a chunk completes, however, the device recurrent state
8624        // owns the sequence and a later decline must be terminal rather than
8625        // falling back to stale CPU state.
8626        let batch_k = std::env::var("CMF_BATCH_K")
8627            .ok()
8628            .and_then(|v| v.parse::<usize>().ok())
8629            .unwrap_or(0);
8630        let batch_admitted = batch_k > 0
8631            && self.can_prefill_batched()
8632            && self.o1_active()
8633            && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
8634            && (0..self.num_layers).all(|li| {
8635                let cache = &self.kv_cache.layers[self.phys_layer(li)];
8636                cache.o1.is_none() || cache.o1_views().is_some()
8637            });
8638        if std::env::var("CMF_GRAPH_PROF").is_ok() {
8639            eprintln!(
8640                "nll-batch: phase=post-seal admission={} requested_k={} scored_rows={}",
8641                batch_admitted,
8642                batch_k,
8643                n.saturating_sub(exact_end),
8644            );
8645        }
8646        let mut batch_completed = false;
8647        if batch_admitted && exact_end < n {
8648            let hs = self.hidden_size;
8649            let mut batch_pos = exact_end;
8650            while batch_pos < n {
8651                let end = (batch_pos + batch_k).min(n);
8652                let bk = end - batch_pos;
8653                let mut hiddens = vec![0.0f32; bk * hs];
8654                for (row, &id) in ids[batch_pos..end].iter().enumerate() {
8655                    hiddens[row * hs..(row + 1) * hs].copy_from_slice(&self.embed_single(id));
8656                }
8657                let positions: Vec<usize> = (batch_pos..end).collect();
8658                let t_batch = std::time::Instant::now();
8659                let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
8660                if std::env::var("CMF_GRAPH_PROF").is_ok() {
8661                    let ms = t_batch.elapsed().as_secs_f64() * 1000.0;
8662                    eprintln!(
8663                        "nll-batch: phase=post-seal mode=o1 k={bk} pos={}..{} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
8664                        batch_pos,
8665                        end.saturating_sub(1),
8666                        bk as f64 / (ms / 1000.0),
8667                    );
8668                }
8669                if let Err(err) = self.nll_check_graph("batch graph", batch_pos) {
8670                    self.nll_end();
8671                    return Err(err);
8672                }
8673                match outcome {
8674                    crate::gpu::BatchGraphOutcome::Completed => {
8675                        batch_completed = true;
8676                        for row in 0..bk {
8677                            nll += self.nll_from_hidden(
8678                                &hiddens[row * hs..(row + 1) * hs],
8679                                ids[batch_pos + row + 1],
8680                                batch_pos + row,
8681                            );
8682                            cnt += 1;
8683                        }
8684                        batch_pos = end;
8685                    }
8686                    crate::gpu::BatchGraphOutcome::Declined => {
8687                        if batch_completed {
8688                            self.nll_end();
8689                            return Err(format!(
8690                                "O(1) NLL batch declined after completed chunk at position {batch_pos}"
8691                            ));
8692                        }
8693                        break;
8694                    }
8695                    crate::gpu::BatchGraphOutcome::Failed => {
8696                        self.nll_end();
8697                        return Err(format!(
8698                            "O(1) NLL batch graph failed after admission at position {batch_pos}"
8699                        ));
8700                    }
8701                }
8702            }
8703            if batch_completed && cnt == n.saturating_sub(requested_start) {
8704                self.nll_end();
8705                return Ok((nll, cnt));
8706            }
8707        }
8708
8709        // Serial O(1) fallback/reference. It is intentionally retained when
8710        // batch admission declines before mutation; callers must label this
8711        // CMF_BATCH_K=0/per-position path separately from the production
8712        // whole-token batch route.
8713        for pos in exact_end..n {
8714            let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8715            if self
8716                .graph_failed
8717                .swap(false, std::sync::atomic::Ordering::Relaxed)
8718            {
8719                self.cancel
8720                    .store(false, std::sync::atomic::Ordering::Relaxed);
8721                self.nll_end();
8722                return Err(format!(
8723                    "GPU graph failed during O(1) NLL serial scoring at position {pos}"
8724                ));
8725            }
8726            nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8727            cnt += 1;
8728        }
8729        self.nll_end();
8730        Ok((nll, cnt))
8731    }
8732
8733    /// Teacher-forced calibration data (B1): for each position, whether the
8734    /// argmax equals the actual next token, and the top-1 softmax prob
8735    /// (top-1 probability) under EACH temperature in `temps` — all from ONE forward
8736    /// pass (argmax/correctness are temperature-invariant; only p_max
8737    /// reshapes). Feeds `cortiq calibrate` (reliability/ECE + temperature
8738    /// fit): is the model's confidence a true property, or does it need a
8739    /// measured scaling?
8740    pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
8741        self.clear_sequence_state();
8742        let n = ids.len().saturating_sub(1);
8743        let mut correct = Vec::with_capacity(n);
8744        let mut pmax = Vec::with_capacity(n);
8745        for pos in 0..n {
8746            let emb = self.embed_single(ids[pos]);
8747            let hidden = self.forward_layers(&emb, pos, None);
8748            let logits = if let Some(logits) = self.graph_logits.take() {
8749                logits
8750            } else {
8751                let normed = inference::rms_norm(
8752                    &hidden,
8753                    &self.weights.final_norm,
8754                    self.rms_eps,
8755                    self.norm_style,
8756                );
8757                // lm_head_forward applies the final-logit softcap itself —
8758                // capping again here double-squashed gemma-class logits
8759                // (tanh∘tanh) and reported a flattered ppl.
8760                self.lm_head_forward(&normed)
8761            };
8762            let target = ids[pos + 1] as usize;
8763            let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
8764            for (i, &v) in logits.iter().enumerate() {
8765                if v > mval {
8766                    mval = v;
8767                    amax = i;
8768                }
8769            }
8770            correct.push(amax == target);
8771            let row: Vec<f32> = temps
8772                .iter()
8773                .map(|&t| {
8774                    let tt = t.max(1e-3);
8775                    let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
8776                    1.0 / s.max(1e-12) // numerator at the max is exp(0)=1
8777                })
8778                .collect();
8779            pmax.push(row);
8780        }
8781        self.clear_sequence_state();
8782        (correct, pmax)
8783    }
8784
8785    /// Teacher-forced PPL with the dynamic router driving per-window
8786    /// skill switches (VMF experiment №2 measurement). Sequential (φ
8787    /// must update per token), returns (ppl, switch_count). The router
8788    /// must be enabled (`enable_dynamic_routing`); else this equals
8789    /// plain `ppl_ids`. The active skill when scoring token t shapes the
8790    /// logits for t+1 — on-policy over the held-out text itself.
8791    pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> Result<(f64, usize), String> {
8792        if self.dyn_router.is_none() {
8793            return Ok((self.ppl_ids(ids)?, 0));
8794        }
8795        self.nll_begin()?;
8796        let saved_active = self.dyn_active;
8797        let mut router = self
8798            .dyn_router
8799            .take()
8800            .ok_or_else(|| "dynamic router disappeared before PPL scoring".to_string())?;
8801        router.reset();
8802        self.dyn_phi_seen = 0;
8803        let _ = self.set_active_skill(None);
8804
8805        let result: Result<(f64, usize), String> = (|| {
8806            let mut nll = 0f64;
8807            let mut cnt = 0usize;
8808            for pos in 0..ids.len().saturating_sub(1) {
8809                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8810                self.nll_check_graph("dynamic serial forward", pos)?;
8811                let out_of_band = self.graph_logits.take();
8812                let mut logits = match out_of_band {
8813                    Some(lg) => lg,
8814                    None => {
8815                        let normed = inference::rms_norm(
8816                            &hidden,
8817                            &self.weights.final_norm,
8818                            self.rms_eps,
8819                            self.norm_style,
8820                        );
8821                        // lm_head_forward applies the final-logit softcap itself —
8822                        // capping again here double-squashed gemma-class logits
8823                        // and reported a flattered ppl.
8824                        self.lm_head_forward(&normed)
8825                    }
8826                };
8827                let target = ids[pos + 1] as usize;
8828                let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8829                let lse: f64 = logits
8830                    .iter()
8831                    .map(|&v| ((v - max) as f64).exp())
8832                    .sum::<f64>()
8833                    .ln()
8834                    + max as f64;
8835                let tok_nll = lse - logits[target] as f64;
8836                if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8837                    let top = logits
8838                        .iter()
8839                        .enumerate()
8840                        .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8841                        .map(|(i, _)| i)
8842                        .unwrap_or(0);
8843                    eprintln!(
8844                        "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8845                        logits[target], logits[top]
8846                    );
8847                }
8848                nll += tok_nll;
8849                cnt += 1;
8850                attention::recycle_buf(&mut logits);
8851                // Route on the evolving phi (drives the NEXT token's skill).
8852                let phi = self.dyn_phi_ema.clone();
8853                if let Some(new_active) = router.step(&phi, pos) {
8854                    let _ = self.set_active_skill(new_active);
8855                }
8856            }
8857            Ok(((nll / cnt.max(1) as f64).exp(), router.switches.len()))
8858        })();
8859
8860        // Restore the detached router and the active overlay on both success
8861        // and failure. The scoring state is cleared independently below.
8862        let _ = self.set_active_skill(saved_active);
8863        self.dyn_router = Some(router);
8864        self.nll_end();
8865        result
8866    }
8867
8868    /// Routing probe φ (spec §9): mean-pooled hidden after `layer`.
8869    pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
8870        self.clear_sequence_state();
8871        let mut acc = vec![0f32; self.hidden_size];
8872        for (pos, &id) in ids.iter().enumerate() {
8873            let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8874            for (a, v) in acc.iter_mut().zip(&h) {
8875                *a += v;
8876            }
8877        }
8878        let n = ids.len().max(1) as f32;
8879        for a in acc.iter_mut() {
8880            *a /= n;
8881        }
8882        self.clear_sequence_state();
8883        acc
8884    }
8885
8886    /// Router-v2 φ probe (spec §9.4, `phi.pool = "span_mean"`): the hidden
8887    /// AFTER `layer` — the same per-position walk and the same quantity as
8888    /// [`Self::probe_phi`] — averaged over the positions in `span` only
8889    /// (the user text between the template's prefix and suffix ids), NOT
8890    /// unit-normalized (the decision normalizes). The walk stops at
8891    /// `span.end`: causality makes the later positions irrelevant, so the
8892    /// result is bit-identical to probing `ids[..span.end]`. An empty span
8893    /// gives the zero vector, which the decision treats as degenerate.
8894    ///
8895    /// Every sequence state is reset before and after — the host KV/ring/
8896    /// recurrent state, the reuse keys (`kv_history`, `kv_prefix`) and the
8897    /// device graph's sequence — so run it on a pipeline that does not
8898    /// also serve a conversation (its prefix reuse would be lost).
8899    pub fn probe_phi_span(
8900        &mut self,
8901        ids: &[u32],
8902        layer: usize,
8903        span: std::ops::Range<usize>,
8904    ) -> Vec<f32> {
8905        #[cfg(target_os = "macos")]
8906        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8907        let end = span.end.min(ids.len());
8908        let start = span.start.min(end);
8909        let reset = |p: &mut Self| p.clear_sequence_state();
8910        reset(self);
8911        let mut acc = vec![0f32; self.hidden_size];
8912        for (pos, &id) in ids[..end].iter().enumerate() {
8913            let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8914            if pos >= start {
8915                for (a, v) in acc.iter_mut().zip(&h) {
8916                    *a += v;
8917                }
8918            }
8919        }
8920        let n = end - start;
8921        if n > 0 {
8922            let n = n as f32;
8923            for a in acc.iter_mut() {
8924                *a /= n;
8925            }
8926        }
8927        reset(self);
8928        acc
8929    }
8930
8931    /// One decode step of the current sequence: forward `token` at
8932    /// `position` (the cache holds positions `[0, position)`, e.g. after
8933    /// [`Self::forward_ids`]) and return the next-token logits — the same
8934    /// forward and head the generation loop runs (resident-graph logits
8935    /// when the graph ran, final norm + lm_head otherwise). The logit-dump
8936    /// tools drive greedy decoding with it so every position is observable.
8937    pub fn decode_step_logits(&mut self, token: u32, position: usize) -> Vec<f32> {
8938        #[cfg(target_os = "macos")]
8939        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8940        self.graph_logits = None;
8941        let hidden = self.forward_layers(&self.embed_single(token), position, None);
8942        if let Some(logits) = self.graph_logits.take() {
8943            return logits;
8944        }
8945        inference::rms_norm_into(
8946            &hidden,
8947            &self.weights.final_norm,
8948            self.rms_eps,
8949            self.norm_style,
8950            &mut self.ws.n1,
8951        );
8952        self.lm_head_forward(&self.ws.n1)
8953    }
8954
8955    /// Layer-major batched prefill (prefill-GEMM): full-attention —
8956    /// per-position with the existing operators (KV grows naturally,
8957    /// causality preserved), GDN projections / FFN / MoE — batched
8958    /// (a weight row is read from DRAM once per chunk, not per
8959    /// position). Returns the hidden of all positions [b × hidden].
8960    fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
8961        self.prefill_batch_masked(ids, start_pos, None)
8962    }
8963
8964    /// `prefill_batch` with a task mask honored on the dense-FFN panels
8965    /// (the masked-inference fast path: full fused compute, mask lands on
8966    /// the activations). The whole-chunk GPU graph is skipped for masked
8967    /// layers by the callers' arms; the per-GEMM device paths stay in
8968    /// play because the zeroing happens on the host between them.
8969    fn prefill_batch_masked(
8970        &mut self,
8971        ids: &[u32],
8972        start_pos: usize,
8973        task_mask: Option<&TaskMask>,
8974    ) -> Vec<f32> {
8975        self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
8976    }
8977
8978    /// One prompt chunk through the whole stack, post-stack rows out (no
8979    /// final norm) — the ingest generation uses, shared by scoring and
8980    /// `forward_ids` so they measure the same execution: the batched wgpu
8981    /// graph's device prefix plus the host's batched walk for the rest when
8982    /// `batch_prefix_prefill` holds and the graph admits the chunk, else
8983    /// the host's chunked prefill. Err only when a graph that had mutated
8984    /// device state failed.
8985    fn prefill_rows(
8986        &mut self,
8987        ids: &[u32],
8988        pos: usize,
8989        task_mask: Option<&TaskMask>,
8990    ) -> Result<Vec<f32>, String> {
8991        self.prefill_input_rows(PrefillIn::Ids(ids), pos, task_mask)
8992    }
8993
8994    fn prefill_input_rows(
8995        &mut self,
8996        input: PrefillIn<'_>,
8997        pos: usize,
8998        task_mask: Option<&TaskMask>,
8999    ) -> Result<Vec<f32>, String> {
9000        self.mimo_moe_prepare();
9001        let hs = self.hidden_size;
9002        let bk = match input {
9003            PrefillIn::Ids(ids) => ids.len(),
9004            PrefillIn::Hidden(rows) => rows.len() / hs,
9005        };
9006        #[cfg(not(target_os = "macos"))]
9007        if task_mask.is_none()
9008            && !self.o1_active()
9009            && bk > 1
9010            && (self.batch_prefix_prefill()
9011                || (self.verify_exact_moe
9012                    && crate::gpu::enabled_here()
9013                    && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)))
9014        {
9015            let mut hiddens = match input {
9016                PrefillIn::Hidden(rows) => rows.to_vec(),
9017                PrefillIn::Ids(ids) => ids.iter().flat_map(|&id| self.embed_single(id)).collect(),
9018            };
9019            let positions: Vec<usize> = (pos..pos + bk).collect();
9020            let mut run = 0usize;
9021            match self.try_batch_graph_wgpu_prefix(
9022                &mut hiddens,
9023                &positions,
9024                bk,
9025                None,
9026                Some(&mut run),
9027            ) {
9028                crate::gpu::BatchGraphOutcome::Completed => {
9029                    let out = if run < self.num_layers {
9030                        self.prefill_batch_span(
9031                            PrefillIn::Hidden(&hiddens),
9032                            pos,
9033                            None,
9034                            run,
9035                            self.num_layers,
9036                        )
9037                    } else {
9038                        hiddens
9039                    };
9040                    return if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
9041                        Err("MiMo attention graph failed after admission".into())
9042                    } else { Ok(out) };
9043                }
9044                crate::gpu::BatchGraphOutcome::Failed => {
9045                    return Err("batched prefix prefill failed after admission".into());
9046                }
9047                crate::gpu::BatchGraphOutcome::Declined => {
9048                    // Rows an earlier chunk left on the device only.
9049                    #[cfg(feature = "gpu")]
9050                    self.pull_lagging_host_kv(0, self.num_layers, pos);
9051                }
9052            }
9053        }
9054        let out = self.prefill_batch_span(input, pos, task_mask, 0, usize::MAX);
9055        if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
9056            Err("batch tail graph failed after admission".into())
9057        } else { Ok(out) }
9058    }
9059
9060    /// The layer-major batched walk over a layer span [from..upto_excl):
9061    /// the whole prefill machinery (chunk graph, batched attends, GEMM
9062    /// panels) for a PARTIAL stack — the network split's prefill rides
9063    /// the same canon as the local one. Input is token ids (embeds
9064    /// itself, coordinator side) or ready boundary hiddens (worker side).
9065    fn prefill_batch_span(
9066        &mut self,
9067        input: PrefillIn<'_>,
9068        start_pos: usize,
9069        task_mask: Option<&TaskMask>,
9070        from: usize,
9071        upto_excl: usize,
9072    ) -> Vec<f32> {
9073        let hs = self.hidden_size;
9074        let b = match input {
9075            PrefillIn::Ids(ids) => ids.len(),
9076            PrefillIn::Hidden(hb) => hb.len() / hs,
9077        };
9078        let upto_excl = upto_excl.min(self.num_layers);
9079        // The CPU embed is deferred: when the chunk graph takes the run
9080        // from layer 0 it gathers the embeddings on the device instead.
9081        // A hidden input is ready by definition.
9082        let mut h: Vec<f32>;
9083        let mut h_ready;
9084        match input {
9085            PrefillIn::Ids(_) => {
9086                h = vec![0.0; b * hs];
9087                h_ready = false;
9088            }
9089            PrefillIn::Hidden(hb) => {
9090                h = hb.to_vec();
9091                h_ready = true;
9092            }
9093        }
9094        let fill_h = |h: &mut Vec<f32>, me: &Self| {
9095            if let PrefillIn::Ids(ids) = input {
9096                for (bi, &id) in ids.iter().enumerate() {
9097                    let e = me.embed_single(id);
9098                    h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
9099                }
9100                if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9101                    if let Ok(t) = tp.parse::<usize>() {
9102                        if t >= start_pos && t < start_pos + ids.len() {
9103                            let bi = t - start_pos;
9104                            let row = &h[bi * hs..(bi + 1) * hs];
9105                            let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9106                            eprintln!(
9107                                "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
9108                                ids[bi],
9109                                row[0],
9110                                row[1],
9111                                ids.len(),
9112                                &ids[..ids.len().min(8)]
9113                            );
9114                        }
9115                    }
9116                }
9117            }
9118        };
9119        let (_nkv, _hd, _rd, eps) = (
9120            self.num_kv_heads,
9121            self.head_dim,
9122            self.rotary_dim,
9123            self.rms_eps,
9124        );
9125        let pool = self.pool.clone();
9126        let norm_style = self.norm_style;
9127        self.mimo_moe_prepare();
9128        let automatic_gpu_prefix = self.automatic_gpu_prefix();
9129
9130        #[cfg(target_os = "macos")]
9131        let mut chunk_skip_until = 0usize;
9132        for li in from..upto_excl {
9133            let _capacity_tail = automatic_gpu_prefix
9134                .filter(|&prefix| {
9135                    li >= prefix && !(self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false))
9136                })
9137                .map(|_| crate::gpu::enter_cpu_scope());
9138            crate::gpu::set_layer(li as i64); // layer-split GPU/CPU
9139            // GPU chunk graph (default-on under CMF_GPU=1): a run of
9140            // consecutive eligible layers for the whole chunk in ONE
9141            // Metal submission — norm, QKV, RoPE with fused mirror
9142            // append, causal attend, O, FFN, hidden device-resident
9143            // across the run. Any refusal falls through to the CPU path.
9144            #[cfg(target_os = "macos")]
9145            if task_mask.is_none() {
9146                if li < chunk_skip_until {
9147                    continue;
9148                }
9149                // Device-side embedding needs a q8_row embedding matrix;
9150                // with any other layout the CPU fills `h` first and the
9151                // graph starts from a ready hidden (refusing the whole
9152                // run over the embedding alone kept q4t models — the
9153                // whole Nanbeige/Bonsai class — on the CPU prefill).
9154                if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
9155                    fill_h(&mut h, self);
9156                    h_ready = true;
9157                }
9158                let ids_for_embed = match input {
9159                    PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
9160                    PrefillIn::Hidden(_) => None,
9161                };
9162                let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
9163                if end > li {
9164                    h_ready = true;
9165                    chunk_skip_until = end;
9166                    // Looped Transformer: the graph stopped at a loop
9167                    // boundary — apply final norm before the next iteration.
9168                    if self.is_loop_end(end - 1) && end < self.num_layers {
9169                        for bi in 0..b {
9170                            let normed = inference::rms_norm(
9171                                &h[bi * hs..(bi + 1) * hs],
9172                                &self.weights.final_norm,
9173                                eps,
9174                                norm_style,
9175                            );
9176                            h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9177                        }
9178                    }
9179                    continue;
9180                }
9181            }
9182            if !h_ready {
9183                fill_h(&mut h, self);
9184                h_ready = true;
9185            }
9186            if task_mask.is_none() && self.verify_exact_moe {
9187                let positions: Vec<_> = (start_pos..start_pos + b).collect();
9188                match self.mimo_graph_layer_rows(li, &mut h, &positions) {
9189                    crate::gpu::BatchGraphOutcome::Completed => continue,
9190                    crate::gpu::BatchGraphOutcome::Failed => return h,
9191                    crate::gpu::BatchGraphOutcome::Declined => {},
9192                }
9193            }
9194            #[cfg(feature = "gpu")]
9195            self.pull_lagging_host_kv(li, li + 1, start_pos);
9196            let lw = &self.weights.layers[self.phys_layer(li)];
9197            let t_attn = std::time::Instant::now();
9198            // ── attention ──
9199            match &lw.attn {
9200                AttnKind::Kda(w) => {
9201                    // Projections batched, recurrence sequential.
9202                    let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
9203                    let mut normed = vec![0.0f32; b * hs];
9204                    for bi in 0..b {
9205                        inference::rms_norm_into(
9206                            &h[bi * hs..(bi + 1) * hs],
9207                            &lw.input_norm,
9208                            eps,
9209                            norm_style,
9210                            &mut normed[bi * hs..(bi + 1) * hs],
9211                        );
9212                    }
9213                    let attn = crate::linear_core::kda_forward_batch(
9214                        &normed,
9215                        b,
9216                        w,
9217                        &cfg,
9218                        &mut self.kv_cache.layers[li].linear_state,
9219                        pool.as_deref(),
9220                    );
9221                    for (dst, &a) in h.iter_mut().zip(&attn) {
9222                        *dst += a;
9223                    }
9224                }
9225                AttnKind::LinearGdn(w) => {
9226                    // Projections batched, recurrence sequential.
9227                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
9228                    let mut normed = vec![0.0f32; b * hs];
9229                    for bi in 0..b {
9230                        let r = inference::rms_norm(
9231                            &h[bi * hs..(bi + 1) * hs],
9232                            &lw.input_norm,
9233                            eps,
9234                            norm_style,
9235                        );
9236                        normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9237                    }
9238                    let attn = crate::linear_core::gdn_forward_batch(
9239                        &normed,
9240                        b,
9241                        w,
9242                        &cfg,
9243                        &mut self.kv_cache.layers[li].linear_state,
9244                        pool.as_deref(),
9245                    );
9246                    for (dst, &a) in h.iter_mut().zip(&attn) {
9247                        *dst += a;
9248                    }
9249                }
9250                AttnKind::ShortConv(w) => {
9251                    // Projections batched over the chunk; the conv walks the
9252                    // contiguous positions in order (same ring as decode).
9253                    let cfg = self
9254                        .short_conv_cfg
9255                        .expect("short-conv layer without short_conv_cfg");
9256                    let mut normed = vec![0.0f32; b * hs];
9257                    for bi in 0..b {
9258                        inference::rms_norm_into(
9259                            &h[bi * hs..(bi + 1) * hs],
9260                            &lw.input_norm,
9261                            eps,
9262                            norm_style,
9263                            &mut normed[bi * hs..(bi + 1) * hs],
9264                        );
9265                    }
9266                    let attn = short_conv_forward_batch(
9267                        &normed,
9268                        b,
9269                        w,
9270                        &cfg,
9271                        &mut self.kv_cache.layers[li].linear_state,
9272                        pool.as_deref(),
9273                    );
9274                    for (dst, &a) in h.iter_mut().zip(&attn) {
9275                        *dst += a;
9276                    }
9277                }
9278                AttnKind::Mla(w) => {
9279                    // Per-position prefill (correctness first; latent
9280                    // batching is a later optimization).
9281                    let inv_freq_l = self.layer_inv_freq(li);
9282                    let rs = self.layer_rope_scale(li);
9283                    let mut normed = vec![0.0f32; hs];
9284                    for bi in 0..b {
9285                        inference::rms_norm_into(
9286                            &h[bi * hs..(bi + 1) * hs],
9287                            &lw.input_norm,
9288                            eps,
9289                            norm_style,
9290                            &mut normed,
9291                        );
9292                        let ao = mla_attention(
9293                            w,
9294                            &normed,
9295                            &mut self.kv_cache.layers[li],
9296                            start_pos + bi,
9297                            &inv_freq_l,
9298                            rs,
9299                            eps,
9300                            pool.as_deref(),
9301                        );
9302                        for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
9303                            *dst += a;
9304                        }
9305                    }
9306                }
9307                AttnKind::Full {
9308                    wq,
9309                    wk,
9310                    wv,
9311                    wo,
9312                    q_norm,
9313                    k_norm,
9314                    output_gate,
9315                    softplus_gate,
9316                    bias,
9317                } => {
9318                    // Chunk-GEMM QKV/O; per-position causal attention
9319                    // inside (roadmap §3 P0 — full-attention prefill no
9320                    // longer re-reads the projection weights b times).
9321                    let mut normed = vec![0.0f32; b * hs];
9322                    for bi in 0..b {
9323                        inference::rms_norm_into(
9324                            &h[bi * hs..(bi + 1) * hs],
9325                            &lw.input_norm,
9326                            eps,
9327                            norm_style,
9328                            &mut normed[bi * hs..(bi + 1) * hs],
9329                        );
9330                    }
9331                    let inv_freq_l = self.layer_inv_freq(li);
9332                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
9333                    let cfg = QwenAttnCfg {
9334                        num_heads: self.layer_num_heads(li),
9335                        num_kv_heads: nkv_l,
9336                        head_dim: hd_l,
9337                        hidden_size: hs,
9338                        position: start_pos,
9339                        inv_freq: &inv_freq_l,
9340                        rotary_dim: rd_l,
9341                        scale: self.attn_scale,
9342                        softcap: self.attn_softcap,
9343                        window: self.layer_window(li),
9344                        v_norm: self.attn_v_norm,
9345                        qk_norm_after_rope: self.qk_norm_after_rope,
9346                        gate_sigmoid: self.proj_gate_sigmoid,
9347                        q_norm: q_norm.as_deref(),
9348                        k_norm: k_norm.as_deref(),
9349                        output_gate: *output_gate,
9350                        softplus_gate: softplus_gate
9351                            .as_ref()
9352                            .map(|(gate, per_head)| (gate, *per_head)),
9353                        rope_scale: self.layer_rope_scale(li),
9354                        bias: bias
9355                            .as_ref()
9356                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
9357                        rms_eps: eps,
9358                        norm_style,
9359                        pool: pool.as_deref(),
9360                        v_head_dim: self.layer_v_dim(li),
9361                    };
9362                    let mut attn = attention::qwen_attention_batch(
9363                        &normed,
9364                        b,
9365                        wq,
9366                        wk,
9367                        wv,
9368                        wo,
9369                        &mut self.kv_cache.layers[li],
9370                        &cfg,
9371                    );
9372                    if let Some(w) = &lw.attn_out_norm {
9373                        for bi in 0..b {
9374                            inference::rms_norm_into(
9375                                &attn[bi * hs..(bi + 1) * hs],
9376                                w,
9377                                eps,
9378                                norm_style,
9379                                &mut normed[bi * hs..(bi + 1) * hs],
9380                            );
9381                        }
9382                        attn.copy_from_slice(&normed);
9383                    }
9384                    for (dst, &a) in h.iter_mut().zip(&attn) {
9385                        *dst += a;
9386                    }
9387                }
9388                AttnKind::Bounded(w) => {
9389                    // Chunk-GEMM projections, the bounded operator per
9390                    // position over ring + chunk — never a growing KV.
9391                    let mut normed = vec![0.0f32; b * hs];
9392                    for bi in 0..b {
9393                        inference::rms_norm_into(
9394                            &h[bi * hs..(bi + 1) * hs],
9395                            &lw.input_norm,
9396                            eps,
9397                            norm_style,
9398                            &mut normed[bi * hs..(bi + 1) * hs],
9399                        );
9400                    }
9401                    let rope = self
9402                        .bounded_rope
9403                        .clone()
9404                        .expect("bounded layer without an installed rotation table");
9405                    let cfg = crate::bounded::BoundedAttnCfg {
9406                        num_heads: self.num_heads,
9407                        num_kv_heads: self.num_kv_heads,
9408                        head_dim: self.head_dim,
9409                        hidden_size: hs,
9410                        scale: self.attn_scale,
9411                        rope: &rope,
9412                        pool: pool.as_deref(),
9413                    };
9414                    let mut attn = crate::bounded::bounded_attention_batch(
9415                        &normed,
9416                        b,
9417                        w,
9418                        &mut self.kv_cache.layers[li],
9419                        &cfg,
9420                    );
9421                    if let Some(wn) = &lw.attn_out_norm {
9422                        for bi in 0..b {
9423                            inference::rms_norm_into(
9424                                &attn[bi * hs..(bi + 1) * hs],
9425                                wn,
9426                                eps,
9427                                norm_style,
9428                                &mut normed[bi * hs..(bi + 1) * hs],
9429                            );
9430                        }
9431                        attn.copy_from_slice(&normed);
9432                    }
9433                    for (dst, &a) in h.iter_mut().zip(&attn) {
9434                        *dst += a;
9435                    }
9436                    attention::recycle_buf(&mut attn);
9437                }
9438                AttnKind::Linear(w) => {
9439                    for bi in 0..b {
9440                        let normed = inference::rms_norm(
9441                            &h[bi * hs..(bi + 1) * hs],
9442                            &lw.input_norm,
9443                            eps,
9444                            norm_style,
9445                        );
9446                        vmf_phase_forward(
9447                            &normed,
9448                            w,
9449                            &self.vmf_cfg.expect("linear layer without vmf_cfg"),
9450                            &mut self.kv_cache.layers[li].linear_state,
9451                            pool.as_deref(),
9452                        )
9453                        .iter()
9454                        .enumerate()
9455                        .for_each(|(i, &a)| h[bi * hs + i] += a);
9456                    }
9457                }
9458            }
9459
9460            // ── FFN batched ──
9461            let lw = &self.weights.layers[self.phys_layer(li)];
9462            let mut post = vec![0.0f32; b * hs];
9463            for bi in 0..b {
9464                let r =
9465                    inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
9466                post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9467            }
9468            // A restrictive per-visit FFN row lands on the activations
9469            // inside the dense arm; an all-open row costs nothing.
9470            let mask_row = task_mask
9471                .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
9472                .and_then(|m| m.ffn_masks.get(li))
9473                .map(|v| v.as_slice());
9474            let attn_ns = t_attn.elapsed().as_nanos() as u64;
9475            let t_ffn = std::time::Instant::now();
9476            let mut ffn = match &lw.ffn {
9477                FfnKind::Dense(d) if !d.segs.is_empty() => {
9478                    tube_ffn(d, &post, b, pool.as_deref(), mask_row)
9479                }
9480                FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
9481                FfnKind::Moe(m) if self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false) => {
9482                    moe_ffn_banked_rows(&mut self.mimo_moe, li, m, &post, b, hs, pool.as_deref())
9483                }
9484                FfnKind::Moe(m) if self.verify_exact_moe => {
9485                    moe_ffn_rows_exact(m, &post, b, hs, pool.as_deref())
9486                }
9487                // Keep prompt expert panels off the projection arena and
9488                // use their routes to prime the model-wide bank.
9489                FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => {
9490                    let before = m.stats.borrow().clone();
9491                    let out = crate::gpu::cpu_scope(|| {
9492                        moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None)
9493                    });
9494                    self.mimo_moe.prime(li, m, &before);
9495                    out
9496                }
9497                FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
9498                // Dual-branch layers run per position (the expert branch
9499                // reads the raw residual — nothing to batch yet).
9500                FfnKind::DenseMoe(dm) => {
9501                    let mut out = vec![0.0f32; b * hs];
9502                    for bi in 0..b {
9503                        let r = dense_moe_ffn(
9504                            dm,
9505                            &post[bi * hs..(bi + 1) * hs],
9506                            &h[bi * hs..(bi + 1) * hs],
9507                            eps,
9508                            norm_style,
9509                            pool.as_deref(),
9510                        );
9511                        out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9512                    }
9513                    out
9514                }
9515            };
9516            if prefill_prof_on() {
9517                PREFILL_SPLIT[0].fetch_add(attn_ns, std::sync::atomic::Ordering::Relaxed);
9518                PREFILL_SPLIT[1].fetch_add(t_ffn.elapsed().as_nanos() as u64, std::sync::atomic::Ordering::Relaxed);
9519            }
9520            if let Some(w) = &lw.ffn_out_norm {
9521                for bi in 0..b {
9522                    inference::rms_norm_into(
9523                        &ffn[bi * hs..(bi + 1) * hs],
9524                        w,
9525                        eps,
9526                        norm_style,
9527                        &mut post[bi * hs..(bi + 1) * hs],
9528                    );
9529                }
9530                ffn.copy_from_slice(&post);
9531            }
9532            for (dst, &f) in h.iter_mut().zip(&ffn) {
9533                *dst += f;
9534            }
9535            if let Some(sc) = lw.layer_scale {
9536                for v in h.iter_mut() {
9537                    *v *= sc;
9538                }
9539            }
9540            // CMF_LAYER_DUMP: every position's hidden after layer li.
9541            if self.layer_dump.is_some() {
9542                for bi in 0..b {
9543                    self.dump_layer_row(start_pos + bi, li, &h[bi * hs..(bi + 1) * hs]);
9544                }
9545            }
9546            if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9547                if let Ok(t) = tp.parse::<usize>() {
9548                    if t >= start_pos && t < start_pos + b {
9549                        let bi = t - start_pos;
9550                        let row = &h[bi * hs..(bi + 1) * hs];
9551                        let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9552                        eprintln!(
9553                            "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
9554                            row[0], row[1]
9555                        );
9556                    }
9557                }
9558            }
9559            // CMF_DEBUG_LAYERS=1: per-layer hidden-state health of the
9560            // LAST prompt position — the knife for "which layer type
9561            // breaks first" on a new architecture.
9562            if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
9563                let row = &h[(b - 1) * hs..b * hs];
9564                let rms =
9565                    (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
9566                let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
9567                eprintln!(
9568                    "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
9569                    match &self.weights.layers[self.phys_layer(li)].attn {
9570                        AttnKind::LinearGdn(_) => "gdn",
9571                        AttnKind::Linear(_) => "vmf",
9572                        AttnKind::ShortConv(_) => "conv",
9573                        _ => "attn",
9574                    },
9575                    match &lw.ffn {
9576                        FfnKind::Moe(_) => "moe",
9577                        FfnKind::Dense(_) => "dense",
9578                        FfnKind::DenseMoe(_) => "dense+moe",
9579                    },
9580                );
9581            }
9582            // Looped Transformer: apply final norm at the end of each loop iteration.
9583            if self.is_loop_end(li) && li + 1 < self.num_layers {
9584                for bi in 0..b {
9585                    let normed = inference::rms_norm(
9586                        &h[bi * hs..(bi + 1) * hs],
9587                        &self.weights.final_norm,
9588                        eps,
9589                        norm_style,
9590                    );
9591                    h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9592                }
9593            }
9594            if std::env::var("CMF_TRACE_H").is_ok() {
9595                let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
9596                let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
9597                eprintln!(
9598                    "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
9599                    lw.layer_scale
9600                );
9601            }
9602        }
9603        crate::gpu::set_layer(-1); // lm_head/final ops outside layer-split
9604        if prefill_prof_on() {
9605            eprintln!(
9606                "prefill-split: attention {:.1} ms, ffn {:.1} ms (cumulative)",
9607                PREFILL_SPLIT[0].load(std::sync::atomic::Ordering::Relaxed) as f64 / 1e6,
9608                PREFILL_SPLIT[1].load(std::sync::atomic::Ordering::Relaxed) as f64 / 1e6
9609            );
9610        }
9611        // A batched span owns a complete set of positions. Publish any
9612        // collecting→sealed transition only after every layer has finished;
9613        // callers that cross into serial/device work must see the new epoch
9614        // before this function returns.
9615        self.o1_progress();
9616        h
9617    }
9618
9619    /// Embed a single token.
9620    fn embed_single(&self, id: u32) -> Vec<f32> {
9621        let mut out = vec![0.0f32; self.hidden_size];
9622        if (id as usize) < self.weights.embed_tokens.rows() {
9623            self.weights.embed_tokens.row_f32(id as usize, &mut out);
9624        }
9625        if self.embed_multiplier != 1.0 {
9626            for v in out.iter_mut() {
9627                *v *= self.embed_multiplier;
9628            }
9629        }
9630        // DeepSeek-V4's hash layers route by TOKEN ID, so the id has to
9631        // reach the forward. It rides in slot 0 (the forward re-reads the
9632        // real embedding itself from the table).
9633        if self.dsv4.is_some()
9634            || self.dsv41.is_some()
9635            || self.qwen4_exp.is_some()
9636        {
9637            let mut v = vec![0.0f32; self.hidden_size.max(1)];
9638            v[0] = id as f32;
9639            return v;
9640        }
9641        // Gemma-3n: the per-layer-embedding half needs the token ID, so
9642        // it rides appended to the embedding; the g3n forward splits it.
9643        if let Some(b) = &self.g3n {
9644            return b.0.extend_embedding(id, &out, self.pool.as_deref());
9645        }
9646        out
9647    }
9648
9649    /// A run of consecutive prefill layers on the GPU for the whole
9650    /// chunk (default-on under CMF_GPU=1; CMF_GPU_CHUNK=0 disables).
9651    /// Eligibility per layer: q8_row weights, plain full attention
9652    /// (no output gate), F32 KV, no o1/masks/gemma extras. Returns the
9653    /// first layer index NOT processed (== `li0` when the run is empty).
9654    #[cfg(target_os = "macos")]
9655    fn chunk_run_gpu(
9656        &mut self,
9657        li0: usize,
9658        h: &mut [f32],
9659        b: usize,
9660        pos0: usize,
9661        embed_ids: Option<&[u32]>,
9662        cap: usize,
9663    ) -> usize {
9664        // (The old streaming attend needed a depth bound at ~1k; the
9665        // GEMM attention scales like the CPU path and lifted it.)
9666        // CMF_GPU_CHUNK=0 disables the graph.
9667        if !crate::gpu::enabled_here()
9668            || std::env::var("CMF_GPU_CHUNK")
9669                .map(|v| v == "0")
9670                .unwrap_or(false)
9671            || b < 32
9672            // Sliding windows with per-layer RoPE ride the chunk graph
9673            // (`metal_graph_swa`: causal_softmax_win, per-layer tables).
9674            || (self.swa.is_some() && !self.metal_graph_swa())
9675            || self.global_attn.is_some()
9676            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
9677            || (self.graph_attn_decline_reason().is_some() && !self.metal_graph_swa())
9678            // Collection owns the exact Q trace and boundary conversion;
9679            // this chunk graph appends dense KV without feeding that trace.
9680            || self.o1_active()
9681            || self.attn_v_norm
9682            || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
9683        {
9684            return li0;
9685        }
9686        let Some(model) = self.model.clone() else {
9687            return li0;
9688        };
9689        let (nh, nkv, hd, hs) = (
9690            self.num_heads,
9691            self.num_kv_heads,
9692            self.head_dim,
9693            self.hidden_size,
9694        );
9695        // Collect the longest run of consecutive eligible layers.
9696        // Looped Transformer: stop at the loop boundary so the CPU can
9697        // apply loop_final_norm between iterations.
9698        let loop_end = if self.loop_final_norm {
9699            ((li0 / self.physical_layers) + 1) * self.physical_layers
9700        } else {
9701            self.num_layers
9702        };
9703        let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
9704        let mut stored_at: Vec<usize> = Vec::new();
9705        let run_end = self.num_layers.min(loop_end).min(cap);
9706        // Each layer's own RoPE table (Spark-X2.5: sliding and full layers
9707        // differ); the global one for every model without sliding layers.
9708        let tables: Vec<std::sync::Arc<Vec<f32>>> =
9709            (li0..run_end).map(|li| self.layer_inv_freq(li)).collect();
9710        for li in li0..run_end {
9711            let lw = &self.weights.layers[self.phys_layer(li)];
9712            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
9713                break;
9714            }
9715            let AttnKind::Full {
9716                wq,
9717                wk,
9718                wv,
9719                wo,
9720                q_norm,
9721                k_norm,
9722                output_gate: false,
9723                softplus_gate,
9724                bias,
9725            } = &lw.attn
9726            else {
9727                break;
9728            };
9729            // Spark-X2.5's per-head sigmoid g_proj gate, from f32 rows.
9730            let head_gate = match softplus_gate {
9731                None => None,
9732                Some((g, true)) if self.proj_gate_sigmoid => match g.f32_parts() {
9733                    Some((d, r, c)) if r == nh && c == hs => Some(d),
9734                    _ => break,
9735                },
9736                Some(_) => break,
9737            };
9738            let FfnKind::Dense(d) = &lw.ffn else { break };
9739            if !matches!(d.act, Act::Silu | Act::Gelu) || !d.segs.is_empty() {
9740                break;
9741            }
9742            // q8_row (row_scale populated), or q4_tiled / q4tp (row_scale
9743            // empty — their scales are in the payload). Mixing across the
9744            // seven projections of one layer is fine; the encoder branches
9745            // per weight on the tensor's dtype. Anything else refuses.
9746            fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
9747                t.q8_row_parts()
9748                    .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9749                    .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9750            }
9751            let parts = (
9752                cw(wq),
9753                cw(wk),
9754                cw(wv),
9755                cw(wo),
9756                cw(&d.gate_proj),
9757                cw(&d.up_proj),
9758                cw(&d.down_proj),
9759            );
9760            let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
9761            else {
9762                break;
9763            };
9764            let layer = &self.kv_cache.layers[li];
9765            if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
9766                break;
9767            }
9768            stored_at.push(layer.head_len(0));
9769            layers.push(crate::gpu_metal::ChunkLayer {
9770                model: &model,
9771                kv_id: self.graph_kv_id,
9772                layer: li,
9773                wq: pq,
9774                wk: pk,
9775                wv: pv,
9776                wo: po,
9777                gate: pg,
9778                up: pu,
9779                down: pd,
9780                input_norm: &lw.input_norm,
9781                post_norm: &lw.post_norm,
9782                bias: bias
9783                    .as_ref()
9784                    .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
9785                q_norm: q_norm.as_deref(),
9786                k_norm: k_norm.as_deref(),
9787                inv_freq: &tables[li - li0],
9788                rd: self.layer_geom(li).2,
9789                nh,
9790                nkv,
9791                hd,
9792                hs,
9793                inter: d.gate_proj.rows(),
9794                gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
9795                late_qk_norm: self.qk_norm_after_rope,
9796                eps: self.rms_eps as f32,
9797                window: self.layer_window(li),
9798                head_gate,
9799                gelu: d.act == Act::Gelu,
9800            });
9801        }
9802        if layers.is_empty() {
9803            return li0;
9804        }
9805        let row = nkv * hd;
9806        let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
9807            .iter()
9808            .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
9809            .collect();
9810        let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
9811        for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
9812            let li = layers[i].layer;
9813            let layer = &self.kv_cache.layers[li];
9814            io.push(crate::gpu_metal::ChunkIo {
9815                cpu_stored: stored_at[i],
9816                cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
9817                cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
9818                out_k: ok,
9819                out_v: ov,
9820                imp: oi,
9821            });
9822        }
9823        let n_run = layers.len();
9824        let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
9825        // Device-side embedding when the run starts the model and the
9826        // embedding matrix is q8_row-mapped.
9827        let ep = embed_ids.and_then(|ids| {
9828            self.weights
9829                .embed_tokens
9830                .q8_row_parts()
9831                .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
9832                    idx,
9833                    rows,
9834                    row_scale: rs,
9835                    ids,
9836                    mult: self.embed_multiplier,
9837                })
9838        });
9839        if embed_ids.is_some() && ep.is_none() {
9840            return li0;
9841        }
9842        if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
9843            return li0;
9844        }
9845        drop(io);
9846        drop(layers);
9847        // CPU caches stay the owners of record: append the chunk rows
9848        // and bank the importance masses per layer.
9849        for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
9850            let li = li0 + i;
9851            let layer = &mut self.kv_cache.layers[li];
9852            for bi in 0..b {
9853                layer.append(
9854                    &ok[bi * row..(bi + 1) * row],
9855                    &ov[bi * row..(bi + 1) * row],
9856                    &[],
9857                );
9858            }
9859            layer.accumulate_imp(oi);
9860        }
9861        last
9862    }
9863
9864    /// Is layer `li` a sliding-window (local-RoPE) layer? Gemma-3:
9865    /// every `pattern`-th layer is global, the rest are local.
9866    fn layer_is_local(&self, li: usize) -> bool {
9867        if let Some(layers) = &self.sliding_layers {
9868            return layers.get(li).copied().unwrap_or(false);
9869        }
9870        match self.swa {
9871            Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
9872            None => false,
9873        }
9874    }
9875
9876    /// The RoPE table for layer `li` (local layers may have their own;
9877    /// Gemma-4 global layers use the proportional padded table).
9878    fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
9879        if self.layer_is_local(li) {
9880            if let Some(f) = &self.inv_freq_local {
9881                return f.clone();
9882            }
9883        } else if let Some(f) = &self.inv_freq_global {
9884            return f.clone();
9885        }
9886        self.inv_freq.clone()
9887    }
9888
9889    /// The attend window for layer `li` (None = full context).
9890    fn layer_window(&self, li: usize) -> Option<usize> {
9891        self.swa
9892            .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
9893    }
9894
9895    fn layer_num_heads(&self, li: usize) -> usize {
9896        self.attention_heads_per_layer
9897            .as_ref()
9898            .and_then(|v| v.get(li).copied())
9899            .unwrap_or(self.num_heads)
9900    }
9901
9902    fn layer_rope_scale(&self, li: usize) -> f32 {
9903        if self.layer_is_local(li) {
9904            self.rope_scale_local
9905        } else {
9906            self.rope_scale
9907        }
9908    }
9909
9910    /// Attention geometry of layer `li`: (num_kv_heads, head_dim,
9911    /// rotary_dim). Gemma-4 global layers override all three.
9912    fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
9913        if !self.layer_is_local(li) {
9914            if let Some((ghd, gkv)) = self.global_attn {
9915                return (gkv, ghd, ghd);
9916            }
9917        }
9918        (
9919            self.layer_num_kv_heads(li),
9920            self.head_dim,
9921            if self.layer_is_local(li) {
9922                self.rotary_dim_local.unwrap_or(self.rotary_dim)
9923            } else {
9924                self.rotary_dim
9925            },
9926        )
9927    }
9928
9929    /// KV heads of layer `li` (virtual index): the per-layer count when the
9930    /// model has one (MiMo-V2), else the uniform `num_kv_heads`.
9931    fn layer_num_kv_heads(&self, li: usize) -> usize {
9932        self.kv_heads_per_layer
9933            .as_ref()
9934            .and_then(|v| v.get(self.phys_layer(li)).copied())
9935            .unwrap_or(self.num_kv_heads)
9936    }
9937
9938    /// V head width of layer `li` (≤ its head_dim).
9939    fn layer_v_dim(&self, li: usize) -> usize {
9940        let (_, hd, _) = self.layer_geom(li);
9941        self.v_head_dim.unwrap_or(hd).min(hd)
9942    }
9943
9944    /// Install a per-layer KV geometry: KV heads per PHYSICAL layer and/or
9945    /// a V head width narrower than `head_dim` (MiMo-V2). Validates it and
9946    /// reshapes the caches of every layer whose KV head count differs from
9947    /// `num_kv_heads`. The loader and the tests share this one path, so a
9948    /// hand-built pipeline cannot hold a geometry the loader would refuse.
9949    /// Call before the first forward (it drops cached rows of reshaped
9950    /// layers). Refuses combinations whose paths would read it wrong:
9951    /// Gemma-4 global layers and MLA carry their own geometry.
9952    pub fn set_attn_geometry(
9953        &mut self,
9954        kv_heads_per_layer: Option<Vec<usize>>,
9955        v_head_dim: Option<usize>,
9956    ) -> Result<(), String> {
9957        if kv_heads_per_layer.is_some() || v_head_dim.is_some() {
9958            if self.global_attn.is_some() {
9959                return Err(
9960                    "per-layer KV heads / v_head_dim cannot combine with Gemma-4 global \
9961                     attention geometry"
9962                        .into(),
9963                );
9964            }
9965            if self
9966                .weights
9967                .layers
9968                .iter()
9969                .any(|lw| matches!(lw.attn, AttnKind::Mla(_)))
9970            {
9971                return Err("per-layer KV heads / v_head_dim cannot combine with MLA".into());
9972            }
9973        }
9974        if let Some(vd) = v_head_dim {
9975            if vd == 0 || vd > self.head_dim {
9976                return Err(format!(
9977                    "v_head_dim {vd} must be in 1..={} (head_dim)",
9978                    self.head_dim
9979                ));
9980            }
9981        }
9982        if let Some(v) = &kv_heads_per_layer {
9983            if v.len() != self.physical_layers {
9984                return Err(format!(
9985                    "kv_heads_per_layer has {} entries, expected {} layers",
9986                    v.len(),
9987                    self.physical_layers
9988                ));
9989            }
9990            for (li, &nkv) in v.iter().enumerate() {
9991                let is_attn = matches!(
9992                    self.weights.layers.get(li).map(|lw| &lw.attn),
9993                    Some(AttnKind::Full { .. }) | None
9994                );
9995                if !is_attn {
9996                    continue;
9997                }
9998                let nh = self
9999                    .attention_heads_per_layer
10000                    .as_ref()
10001                    .and_then(|h| h.get(li).copied())
10002                    .unwrap_or(self.num_heads);
10003                if nkv == 0 || nh % nkv != 0 {
10004                    return Err(format!(
10005                        "layer {li}: {nkv} KV heads must be nonzero and divide {nh} Q heads"
10006                    ));
10007                }
10008            }
10009        }
10010        self.kv_heads_per_layer = kv_heads_per_layer;
10011        self.v_head_dim = v_head_dim.filter(|&vd| vd != self.head_dim);
10012        if self.kv_heads_per_layer.is_some() {
10013            for li in 0..self.kv_cache.layers.len() {
10014                let full = matches!(
10015                    self.weights
10016                        .layers
10017                        .get(self.phys_layer(li))
10018                        .map(|lw| &lw.attn),
10019                    Some(AttnKind::Full { .. })
10020                );
10021                let nkv = self.layer_num_kv_heads(li);
10022                let cache = &self.kv_cache.layers[li];
10023                if full && (cache.num_kv_heads != nkv || cache.head_dim != self.head_dim) {
10024                    let sinks = cache.sinks.clone();
10025                    self.kv_cache.layers[li] =
10026                        crate::kv_cache::LayerKvCache::new(nkv, self.head_dim);
10027                    self.kv_cache.layers[li].sinks = sinks;
10028                }
10029            }
10030        }
10031        Ok(())
10032    }
10033
10034    /// Attach learned attention-sink logits (one per Q head) to PHYSICAL
10035    /// layer `phys` — every virtual layer that runs it. The loader calls
10036    /// this for each `model.layers.N.self_attn.sinks` tensor.
10037    pub fn set_layer_sinks(&mut self, phys: usize, sinks: Vec<f32>) -> Result<(), String> {
10038        let Some(lw) = self.weights.layers.get(phys) else {
10039            return Err(format!("sinks for layer {phys}: no such layer"));
10040        };
10041        if !matches!(lw.attn, AttnKind::Full { .. }) {
10042            return Err(format!(
10043                "sinks for layer {phys}: only softmax (Full) attention layers take sinks"
10044            ));
10045        }
10046        let nh = self
10047            .attention_heads_per_layer
10048            .as_ref()
10049            .and_then(|h| h.get(phys).copied())
10050            .unwrap_or(self.num_heads);
10051        if sinks.len() != nh {
10052            return Err(format!(
10053                "sinks for layer {phys}: {} values, expected one per Q head ({nh})",
10054                sinks.len()
10055            ));
10056        }
10057        if let Some(bad) = sinks.iter().find(|v| !v.is_finite()) {
10058            return Err(format!("sinks for layer {phys}: non-finite value {bad}"));
10059        }
10060        for li in 0..self.kv_cache.layers.len() {
10061            if self.phys_layer(li) == phys {
10062                self.kv_cache.layers[li].sinks = Some(sinks.clone());
10063            }
10064        }
10065        Ok(())
10066    }
10067
10068    /// Why the GPU attention graphs cannot serve this model, if they
10069    /// cannot: the wgpu whole-token and batched graphs, the greedy
10070    /// multi-burst, the q1 attention dropin and the Metal block/chunk/rows
10071    /// graphs all assume ONE (num_kv_heads, head_dim) geometry, V heads as
10072    /// wide as K, a single RoPE table, full-context attention and a plain
10073    /// softmax. A model outside that contract runs on the CPU layer walk
10074    /// (and the per-op GPU matvecs) until a graph learns it — never on a
10075    /// graph that would read it wrong. None = no attention-level reason
10076    /// (the graph builders still check weights and layer kinds).
10077    pub fn graph_attn_decline_reason(&self) -> Option<&'static str> {
10078        if self.kv_heads_per_layer.is_some() {
10079            return Some("per-layer KV head counts");
10080        }
10081        if self.v_head_dim.is_some_and(|vd| vd != self.head_dim) {
10082            return Some("V heads narrower than Q/K heads");
10083        }
10084        if self.kv_cache.layers.iter().any(|l| l.sinks.is_some()) {
10085            return Some("learned attention sinks");
10086        }
10087        if self.swa.is_some() || self.sliding_layers.is_some() {
10088            return Some("sliding-window layers");
10089        }
10090        None
10091    }
10092
10093    /// Can the Metal block graph carry this model's sliding-window layers?
10094    /// Both of its attention forms take a layer's own window, RoPE table
10095    /// and rotary width: the device attend (`AttnDeviceParams::window`, the
10096    /// layer's `inv_freq`/`rd`) and the sandwich's host attend (exactly the
10097    /// CPU path's attention). True when the sliding window is the model's
10098    /// only attention-level decline (Spark-X2.5); per-layer KV heads,
10099    /// narrow V, sinks, Gemma-4 global geometry and scaled RoPE positions
10100    /// still decline.
10101    ///
10102    /// Per-layer Q heads (Laguna) and capped scores (Gemma-2) decline here
10103    /// too, with V norm (Gemma-4): the chunk prefill's own gate checks
10104    /// neither of the first two, and before this door opened its `swa`
10105    /// refusal kept every sliding model with them off the chunk graph.
10106    #[cfg(target_os = "macos")]
10107    fn metal_graph_swa(&self) -> bool {
10108        (self.swa.is_some() || self.sliding_layers.is_some())
10109            && self.attention_heads_per_layer.is_none()
10110            && !self.attn_v_norm
10111            && self.attn_softcap == 0.0
10112            && self.kv_heads_per_layer.is_none()
10113            && self.v_head_dim.map_or(true, |vd| vd == self.head_dim)
10114            && !self.kv_cache.layers.iter().any(|l| l.sinks.is_some())
10115            && self.global_attn.is_none()
10116            && self.inv_freq_global.is_none()
10117            && (0..self.num_layers).all(|li| self.layer_rope_scale(li) == 1.0)
10118            && !(0..self.num_layers).any(|li| {
10119                self.layer_is_local(li)
10120                    && self.inv_freq_local.is_none()
10121                    && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
10122            })
10123    }
10124
10125    /// Why the WGPU graphs (whole-token, batched prefill, greedy burst)
10126    /// cannot run this model's attention, if they cannot. Per-layer KV
10127    /// heads, V narrower than K, learned sinks and sliding windows ride
10128    /// their per-layer geometry (`GraphAttnGeom`, the ATTEND_X kernels);
10129    /// what that geometry does not express keeps the decline, by name.
10130    /// None for every model with one attention geometry.
10131    pub fn wgpu_graph_attn_decline(&self) -> Option<&'static str> {
10132        self.graph_attn_decline_reason()?;
10133        if self.global_attn.is_some() {
10134            return Some("per-layer head width (Gemma-4 global layers) with per-layer geometry");
10135        }
10136        if self.attention_heads_per_layer.is_some() {
10137            return Some("per-layer Q head counts with per-layer geometry");
10138        }
10139        if self.attn_v_norm {
10140            return Some("V norm with per-layer geometry");
10141        }
10142        if (0..self.num_layers).any(|li| self.layer_rope_scale(li) != 1.0) {
10143            return Some("scaled RoPE positions with per-layer geometry");
10144        }
10145        if self.weights.layers.iter().any(|lw| {
10146            lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some()
10147        }) {
10148            return Some("sandwich norms / layer scale with per-layer geometry");
10149        }
10150        if self.weights.layers.iter().any(|lw| {
10151            matches!(
10152                &lw.attn,
10153                AttnKind::Full {
10154                    output_gate: true,
10155                    ..
10156                }
10157            )
10158        }) && self.v_head_dim.is_some()
10159        {
10160            return Some("gated attention with V narrower than K");
10161        }
10162        if (0..self.num_layers).any(|li| {
10163            self.layer_is_local(li)
10164                && self.inv_freq_local.is_none()
10165                && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
10166        }) {
10167            return Some("local rotary width without a local RoPE table");
10168        }
10169        None
10170    }
10171
10172    /// The wgpu graphs' attention geometry for layer `li` (virtual index):
10173    /// Some only for a model whose layers do not share one (MiMo-V2) — KV
10174    /// heads, V width, rotary width and RoPE table, window and sinks of
10175    /// THIS layer, exactly what the CPU attention reads for it.
10176    fn graph_attn_geom(&self, li: usize) -> Option<crate::gpu::GraphAttnGeom<'_>> {
10177        self.graph_attn_decline_reason()?;
10178        let (nkv, _hd, rd) = self.layer_geom(li);
10179        let invf: &[f32] = if self.layer_is_local(li) {
10180            match &self.inv_freq_local {
10181                Some(f) => f.as_slice(),
10182                None => self.inv_freq.as_slice(),
10183            }
10184        } else {
10185            match &self.inv_freq_global {
10186                Some(f) => f.as_slice(),
10187                None => self.inv_freq.as_slice(),
10188            }
10189        };
10190        Some(crate::gpu::GraphAttnGeom {
10191            nkv,
10192            dv: self.layer_v_dim(li),
10193            rd,
10194            invf,
10195            window: self.layer_window(li),
10196            sink: self.kv_cache.layers[li].sinks.as_deref(),
10197        })
10198    }
10199
10200    /// Bring the host KV cache of every Full-attention layer in
10201    /// `[from, upto)` up to `position` rows from the wgpu mirrors, where a
10202    /// device graph advanced a layer that the host is about to run: a
10203    /// device prefix that shrank since the prompt (or a batched prefill
10204    /// prefix longer than the decode one). A layer whose mirror does not
10205    /// hold the missing rows is left alone. Rows a sliding layer's ring
10206    /// no longer holds come back as zeros — outside every window that
10207    /// will read them.
10208    #[cfg(feature = "gpu")]
10209    fn pull_lagging_host_kv(&mut self, from: usize, upto: usize, position: usize) {
10210        let kv_id = self.graph_kv_id;
10211        for li in from..upto.min(self.num_layers) {
10212            if !matches!(
10213                self.weights.layers[self.phys_layer(li)].attn,
10214                AttnKind::Full { .. }
10215            ) {
10216                continue;
10217            }
10218            let host = self.kv_cache.layers[li].seq_len;
10219            if host >= position {
10220                continue;
10221            }
10222            let Some(dev) = crate::gpu::graph_kv_stored(kv_id, li) else {
10223                continue;
10224            };
10225            let to = dev.min(position);
10226            if to <= host {
10227                continue;
10228            }
10229            let (nkv, hd) = {
10230                let c = &self.kv_cache.layers[li];
10231                (c.num_kv_heads, c.head_dim)
10232            };
10233            let Some((k, v, first_valid)) =
10234                crate::gpu::graph_kv_pull_host(kv_id, li, host, to, nkv, hd)
10235            else {
10236                continue;
10237            };
10238            // A sliding layer only ever reads its last `window` rows; a
10239            // full-context layer needs every row it did not have.
10240            let need_from = match self.layer_window(li) {
10241                Some(w) => host.max((position + 1).saturating_sub(w)),
10242                None => host,
10243            };
10244            if first_valid > need_from {
10245                tracing::warn!(
10246                    "layer {li}: device KV rows {host}..{to} no longer resident \
10247                     (from {first_valid}); host attention will miss them"
10248                );
10249            }
10250            let row = nkv * hd;
10251            let cache = &mut self.kv_cache.layers[li];
10252            for p in 0..to - host {
10253                cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
10254            }
10255        }
10256    }
10257
10258    /// Log (once per graph site and pipeline) that `site` declined for
10259    /// `reason`. The lines are kept so a caller or a test can read them.
10260    fn note_graph_decline(&self, site: &'static str, reason: &'static str) {
10261        let mut seen = self.graph_declines.borrow_mut();
10262        if !seen.iter().any(|&(s, r)| s == site && r == reason) {
10263            tracing::warn!("{site} declined: {reason} (CPU attention path)");
10264            seen.push((site, reason));
10265        }
10266    }
10267
10268    /// The GPU-graph declines this pipeline has logged so far, as the
10269    /// logged lines.
10270    pub fn graph_declines(&self) -> Vec<String> {
10271        self.graph_declines
10272            .borrow()
10273            .iter()
10274            .map(|(site, reason)| format!("{site} declined: {reason} (CPU attention path)"))
10275            .collect()
10276    }
10277
10278    /// Does layer `li` have the plain attention geometry the historical
10279    /// head-masked f32 path (`multi_head_attention`) assumes — pipeline-wide
10280    /// KV heads / head_dim / RoPE table, full context, no sink, V as wide
10281    /// as K? Anything else runs the dense `qwen_attention` instead.
10282    /// `CMF_LAYER_DUMP` writer (see `Pipeline::layer_dump`): one position's
10283    /// hidden after layer `li` as raw little-endian f32 into
10284    /// `<dir>/p{pos:06}_l{li:02}.f32`. A failed write is reported once and
10285    /// never stops the forward — the dump is a diagnostic.
10286    fn dump_layer_row(&self, pos: usize, li: usize, row: &[f32]) {
10287        let Some(dir) = &self.layer_dump else {
10288            return;
10289        };
10290        let mut bytes = Vec::with_capacity(row.len() * 4);
10291        for v in row {
10292            bytes.extend_from_slice(&v.to_le_bytes());
10293        }
10294        let path = dir.join(format!("p{pos:06}_l{li:02}.f32"));
10295        if let Err(e) = std::fs::create_dir_all(dir).and_then(|_| std::fs::write(&path, &bytes)) {
10296            use std::sync::atomic::{AtomicBool, Ordering};
10297            static SAID: AtomicBool = AtomicBool::new(false);
10298            if !SAID.swap(true, Ordering::Relaxed) {
10299                tracing::error!("CMF_LAYER_DUMP: cannot write {}: {e}", path.display());
10300            }
10301        }
10302    }
10303
10304    /// Decide the MiMo-V2 expert placement once (`crate::mimo_moe`). Any
10305    /// other model turns the slot off on the first call.
10306    fn mimo_moe_prepare(&mut self) {
10307        if !self.mimo_moe.is_undecided() {
10308            return;
10309        }
10310        let slot = {
10311            let layers: Vec<(usize, &MoeFfn)> = (0..self.num_layers)
10312                .filter_map(
10313                    |li| match &self.weights.layers.get(self.phys_layer(li))?.ffn {
10314                        FfnKind::Moe(m) => Some((li, m)),
10315                        _ => None,
10316                    },
10317                )
10318                .collect();
10319            // One bank lives on one device: an in-process multi-GPU split
10320            // keeps the whole-layer path.
10321            if layers.is_empty()
10322                || self.physical_layers != self.num_layers
10323                || self.gpu_plan.is_some()
10324            {
10325                crate::mimo_moe::Slot::Off
10326            } else {
10327                // Whether a whole-token graph could run this model's layers
10328                // (then a whole-layer prefix is one submit, not per-layer
10329                // fences).
10330                let graph_prefix =
10331                    self.wgpu_graph_attn_decline().is_none() && crate::gpu::wgpu_graph_default();
10332                crate::mimo_moe::Slot::decide(&layers, self.num_layers, graph_prefix)
10333            }
10334        };
10335        self.mimo_moe = slot;
10336    }
10337
10338    #[cfg(test)]
10339    pub(crate) fn test_graph_kv_id(&self) -> u64 {
10340        self.graph_kv_id
10341    }
10342
10343    /// Dynamic MiMo layer: one device attention graph, followed by a bank
10344    /// frame. Both decode and short verification use this same attention
10345    /// path and absolute layer key; the host KV may intentionally lag.
10346    pub(crate) fn mimo_graph_layer_rows(
10347        &mut self,
10348        li: usize,
10349        h: &mut [f32],
10350        positions: &[usize],
10351    ) -> crate::gpu::BatchGraphOutcome {
10352        use crate::gpu::BatchGraphOutcome as Out;
10353        let b = positions.len();
10354        if !(1..=4).contains(&b)
10355            || h.len() != b * self.hidden_size
10356            || !self.mimo_moe.is_dynamic(li, true)
10357            || !crate::gpu::enabled_here()
10358            || !crate::gpu::wgpu_active()
10359            || self.o1_active()
10360            || self.physical_layers != self.num_layers
10361            // The pair-fusion diagnostic (and an explicit graph-off run)
10362            // rewinds only host KV. A hidden singleton attention graph here
10363            // would leave device mirrors ahead of the next host position.
10364            || std::env::var("CMF_GPU_WGPU_GRAPH").as_deref() == Ok("0")
10365            || std::env::var("CMF_MIMO_ATTN_GRAPH").as_deref() == Ok("0")
10366            || self.wgpu_graph_attn_decline().is_some()
10367        {
10368            return Out::Declined;
10369        }
10370        let attn_started = std::time::Instant::now();
10371        let outcome = {
10372            let lw = &self.weights.layers[li];
10373            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
10374                return Out::Declined;
10375            }
10376            let FfnKind::Moe(m) = &lw.ffn else {
10377                return Out::Declined;
10378            };
10379            let AttnKind::Full {
10380                wq,
10381                wk,
10382                wv,
10383                wo,
10384                q_norm,
10385                k_norm,
10386                output_gate,
10387                softplus_gate,
10388                bias,
10389            } = &lw.attn
10390            else {
10391                return Out::Declined;
10392            };
10393            if *output_gate || softplus_gate.is_some() {
10394                return Out::Declined;
10395            }
10396            let Some(model) = wq.graph_weight().map(|(m, ..)| m.clone()).or_else(|| {
10397                m.experts
10398                    .first()?
10399                    .gate_proj
10400                    .mapped_q4tp()
10401                    .map(|(m, _)| m.clone())
10402            }) else {
10403                return Out::Declined;
10404            };
10405            fn gw<'a>(
10406                t: &'a QTensor,
10407                owner: &std::sync::Arc<cortiq_core::CmfModel>,
10408            ) -> Option<crate::gpu::GraphW<'a>> {
10409                if let Some((m, idx, kind, rs)) = t.graph_weight() {
10410                    if m.uid() != owner.uid() || t.has_prism_contract() {
10411                        return None;
10412                    }
10413                    return Some(crate::gpu::GraphW {
10414                        idx,
10415                        kind,
10416                        row_scale: rs,
10417                        data: &[],
10418                        prism: crate::gpu::GraphPrismOp::None,
10419                        affine: false,
10420                    });
10421                }
10422                t.as_f32().map(|data| crate::gpu::GraphW {
10423                    idx: 0,
10424                    kind: 4,
10425                    row_scale: &[],
10426                    data,
10427                    prism: crate::gpu::GraphPrismOp::None,
10428                    affine: false,
10429                })
10430            }
10431            let (Some(q), Some(k), Some(v), Some(o)) = (
10432                gw(wq, &model),
10433                gw(wk, &model),
10434                gw(wv, &model),
10435                gw(wo, &model),
10436            ) else {
10437                return Out::Declined;
10438            };
10439            let layer = crate::gpu::GraphLayer {
10440                input_norm: &lw.input_norm,
10441                post_norm: &lw.post_norm,
10442                ffn: crate::gpu::GraphFfn::AttentionOnly,
10443                attn: crate::gpu::GraphAttn::Full {
10444                    wq: q,
10445                    wk: k,
10446                    wv: v,
10447                    wo: o,
10448                    q_norm: q_norm.as_deref(),
10449                    k_norm: k_norm.as_deref(),
10450                    late_qk_norm: self.qk_norm_after_rope,
10451                    bias: bias
10452                        .as_ref()
10453                        .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice())),
10454                    output_gate: false,
10455                    cpu_k: self.kv_cache.layers[li].k_heads(),
10456                    cpu_v: self.kv_cache.layers[li].v_heads(),
10457                    geom: self.graph_attn_geom(li),
10458                    head_gate: None,
10459                },
10460            };
10461            let (nkv, hd, rd) = self.layer_geom(li);
10462            crate::gpu::forward_batch_graph_at(
10463                &model,
10464                self.graph_kv_id,
10465                li,
10466                &[layer],
10467                &self.inv_freq,
10468                h,
10469                self.layer_num_heads(li),
10470                nkv,
10471                hd,
10472                rd,
10473                self.hidden_size,
10474                1,
10475                positions,
10476                self.kv_cache.max_seq_len,
10477                self.norm_style == cortiq_core::NormStyle::Gemma,
10478                self.rms_eps as f32,
10479                self.attn_scale,
10480                b,
10481                &[],
10482                self.o1_epoch,
10483                None,
10484                None,
10485            )
10486        };
10487        match outcome {
10488            Out::Completed => {}
10489            Out::Declined => return Out::Declined,
10490            Out::Failed => {
10491                self.graph_failed
10492                    .store(true, std::sync::atomic::Ordering::Relaxed);
10493                return Out::Failed;
10494            }
10495        }
10496        let attn_ns = attn_started.elapsed().as_nanos() as u64;
10497        let hs = self.hidden_size;
10498        let lw = &self.weights.layers[li];
10499        let FfnKind::Moe(m) = &lw.ffn else {
10500            unreachable!()
10501        };
10502        let mut post = vec![0.0; h.len()];
10503        for (x, y) in h.chunks_exact(hs).zip(post.chunks_exact_mut(hs)) {
10504            inference::rms_norm_into(x, &lw.post_norm, self.rms_eps, self.norm_style, y);
10505        }
10506        let mut ffn = if b == 1 {
10507            moe_ffn_banked(&mut self.mimo_moe, li, m, &post, self.pool.as_deref())
10508        } else {
10509            moe_ffn_banked_rows(
10510                &mut self.mimo_moe,
10511                li,
10512                m,
10513                &post,
10514                b,
10515                hs,
10516                self.pool.as_deref(),
10517            )
10518        };
10519        for (x, &f) in h.iter_mut().zip(&ffn) {
10520            *x += f;
10521        }
10522        attention::recycle_buf(&mut ffn);
10523        if self.layer_dump.is_some() {
10524            for (&pos, row) in positions.iter().zip(h.chunks_exact(hs)) {
10525                self.dump_layer_row(pos, li, row);
10526            }
10527        }
10528        crate::mimo_moe::note_attention_graph(b, attn_ns);
10529        Out::Completed
10530    }
10531
10532    fn layer_attn_plain(&self, li: usize) -> bool {
10533        self.kv_heads_per_layer.is_none()
10534            && self.v_head_dim.is_none()
10535            && self.global_attn.is_none()
10536            && self.layer_window(li).is_none()
10537            && self.kv_cache.layers[li].sinks.is_none()
10538    }
10539
10540    /// Forward one position through all layers (hybrid dispatch).
10541    fn forward_layers(
10542        &mut self,
10543        hidden: &[f32],
10544        position: usize,
10545        task_mask: Option<&TaskMask>,
10546    ) -> Vec<f32> {
10547        let out = self.forward_layers_upto(hidden, position, task_mask, None);
10548        self.o1_progress();
10549        out
10550    }
10551
10552    // ── Network pipeline-split building blocks (coordinator/worker) ──
10553    // A remote worker owns layers [from ..= upto] and their KV; the
10554    // coordinator owns the rest plus embed / final norm / head. Attention
10555    // causality is per-layer, so a whole prompt's boundary hiddens ship
10556    // as one batch and decode ships one vector per token.
10557
10558    /// Embed one token id (embed multiplier applied).
10559    pub fn embed_id(&self, id: u32) -> Vec<f32> {
10560        self.embed_single(id)
10561    }
10562
10563    /// Refuse the archs/modes whose forward cannot be cut at a layer
10564    /// boundary. Loud by design: a split that silently changed the math
10565    /// would be a chimera.
10566    pub fn split_supported(&self) -> Result<(), String> {
10567        if self.dsv4.is_some() {
10568            return Err(
10569                "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
10570            );
10571        }
10572        if self.dsv41.is_some() {
10573            return Err(
10574                "network split: DeepSeek-V4.1 owns the shared CED/CSA2 state (not splittable)"
10575                    .into(),
10576            );
10577        }
10578        if self.qwen4_exp.is_some() {
10579            return Err(
10580                "network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
10581            );
10582        }
10583        if self.g3n.is_some() {
10584            return Err(
10585                "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
10586            );
10587        }
10588        Ok(())
10589    }
10590
10591    /// Forward `hidden` through layers [from ..= upto] at `position`,
10592    /// appending those layers' KV/state. Both split sides call this
10593    /// over their own range; a task mask applies to the span's own
10594    /// layers (each side masks what it runs).
10595    pub fn forward_span(
10596        &mut self,
10597        hidden: &[f32],
10598        position: usize,
10599        from: usize,
10600        upto: usize,
10601        task_mask: Option<&TaskMask>,
10602    ) -> Result<Vec<f32>, String> {
10603        self.split_supported()?;
10604        if from > upto || upto >= self.num_layers {
10605            return Err(format!(
10606                "forward_span: layer range {from}..={upto} outside 0..{}",
10607                self.num_layers
10608            ));
10609        }
10610        if hidden.len() != self.hidden_size {
10611            return Err(format!(
10612                "forward_span: hidden len {} ≠ hidden_size {}",
10613                hidden.len(),
10614                self.hidden_size
10615            ));
10616        }
10617        let out = self.forward_layers_span(hidden, position, task_mask, from, Some(upto));
10618        self.o1_progress();
10619        if self
10620            .graph_failed
10621            .swap(false, std::sync::atomic::Ordering::Relaxed)
10622        {
10623            self.cancel
10624                .store(false, std::sync::atomic::Ordering::Relaxed);
10625            self.clear_sequence_state();
10626            return Err("forward_span: deferred O(1) transition failed".into());
10627        }
10628        Ok(out)
10629    }
10630
10631    /// Final norm + lm_head over a boundary hidden (the final-logit
10632    /// softcap is applied by lm_head_forward itself).
10633    pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
10634        let normed = inference::rms_norm(
10635            hidden,
10636            &self.weights.final_norm,
10637            self.rms_eps,
10638            self.norm_style,
10639        );
10640        self.lm_head_forward(&normed)
10641    }
10642
10643    /// Sample the next token with this pipeline's sampler state.
10644    pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
10645        sampler::sample_with_scratch(
10646            logits,
10647            &self.sampler_config,
10648            past_tokens,
10649            &mut self.rng,
10650            &mut self.sampler_scratch,
10651        )
10652    }
10653
10654    /// Fresh sequence: clear KV, reuse history and device mirrors.
10655    pub fn reset_session(&mut self) {
10656        self.clear_sequence_state();
10657    }
10658
10659    /// Batched span prefill from token ids (coordinator side): embed +
10660    /// layers [0 ..= upto]; returns the boundary hiddens of ALL positions
10661    /// (ids.len() × hidden). Rides the same layer-major machinery as the
10662    /// local prefill; falls back to the per-position walk under
10663    /// CMF_PREFILL=seq.
10664    pub fn prefill_span_ids(
10665        &mut self,
10666        ids: &[u32],
10667        start_pos: usize,
10668        upto: usize,
10669        task_mask: Option<&TaskMask>,
10670    ) -> Result<Vec<f32>, String> {
10671        self.split_supported()?;
10672        if upto >= self.num_layers {
10673            return Err(format!(
10674                "prefill_span_ids: upto {upto} outside 0..{}",
10675                self.num_layers
10676            ));
10677        }
10678        // Same predicate as the whole-stack prefill: a span whose GDN
10679        // state lives on the device must walk positions through the
10680        // graph, not through the batched CPU span.
10681        if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10682            let out =
10683                self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1);
10684            self.check_o1_progress_failure("prefill_span_ids")?;
10685            Ok(out)
10686        } else {
10687            let hs = self.hidden_size;
10688            let mut out = Vec::with_capacity(ids.len() * hs);
10689            for (i, &id) in ids.iter().enumerate() {
10690                let emb = self.embed_id(id);
10691                out.extend_from_slice(&self.forward_span(
10692                    &emb,
10693                    start_pos + i,
10694                    0,
10695                    upto,
10696                    task_mask,
10697                )?);
10698            }
10699            Ok(out)
10700        }
10701    }
10702
10703    /// Batched span prefill from boundary hiddens (worker side): layers
10704    /// [from ..= upto] for every position in the batch; returns the batch.
10705    pub fn prefill_span_hidden(
10706        &mut self,
10707        hidden: &[f32],
10708        start_pos: usize,
10709        from: usize,
10710        upto: usize,
10711        task_mask: Option<&TaskMask>,
10712    ) -> Result<Vec<f32>, String> {
10713        self.split_supported()?;
10714        let hs = self.hidden_size;
10715        if hidden.is_empty() || hidden.len() % hs != 0 {
10716            return Err(format!(
10717                "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
10718                hidden.len()
10719            ));
10720        }
10721        if from > upto || upto >= self.num_layers {
10722            return Err(format!(
10723                "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
10724                self.num_layers
10725            ));
10726        }
10727        if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10728            let out = self.prefill_batch_span(
10729                PrefillIn::Hidden(hidden),
10730                start_pos,
10731                task_mask,
10732                from,
10733                upto + 1,
10734            );
10735            self.check_o1_progress_failure("prefill_span_hidden")?;
10736            Ok(out)
10737        } else {
10738            let b = hidden.len() / hs;
10739            let mut out = Vec::with_capacity(hidden.len());
10740            for i in 0..b {
10741                let h = self.forward_span(
10742                    &hidden[i * hs..(i + 1) * hs],
10743                    start_pos + i,
10744                    from,
10745                    upto,
10746                    task_mask,
10747                )?;
10748                out.extend_from_slice(&h);
10749            }
10750            Ok(out)
10751        }
10752    }
10753
10754    /// Build the whole-token wgpu graph for a pure-attention q1 model (every
10755    /// layer Full q1 + dense q1 FFN, no gate/bias). Returns the post-stack
10756    /// hidden (caller does final norm + lm_head), or None to fall back.
10757    fn try_token_graph_wgpu(
10758        &self,
10759        hidden: &[f32],
10760        position: usize,
10761        logits_out: &mut Vec<f32>,
10762        layers_run: &mut usize,
10763    ) -> Option<Result<Vec<f32>, ()>> {
10764        self.try_token_graph_wgpu_steps(
10765            hidden,
10766            position,
10767            logits_out,
10768            1,
10769            None,
10770            Some(layers_run),
10771            0,
10772            self.num_layers,
10773        )
10774    }
10775
10776    /// The span twin (network split): the graph covers [from..upto_excl)
10777    /// — one submit per SEGMENT per token. lm_head folds in only when
10778    /// the span reaches the last layer.
10779    fn try_token_graph_wgpu_span(
10780        &self,
10781        hidden: &[f32],
10782        position: usize,
10783        logits_out: &mut Vec<f32>,
10784        from: usize,
10785        upto_excl: usize,
10786        layers_run: &mut usize,
10787    ) -> Option<Result<Vec<f32>, ()>> {
10788        self.try_token_graph_wgpu_steps(
10789            hidden,
10790            position,
10791            logits_out,
10792            1,
10793            None,
10794            Some(layers_run),
10795            from,
10796            upto_excl,
10797        )
10798    }
10799
10800    /// Greedy burst: forward `t_next` and let the device pick + re-embed
10801    /// the next k−1 tokens — k frames, ONE submit, k ids back. The ZML
10802    /// trade, on wgpu. None ⇒ caller keeps the per-token path.
10803    fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
10804        if self.o1_active() || self.attn_softcap > 0.0 {
10805            return None;
10806        }
10807        // The burst builds the whole-token graph; attention the graph's
10808        // per-layer geometry cannot express keeps the per-token path.
10809        if let Some(reason) = self.wgpu_graph_attn_decline() {
10810            self.note_graph_decline("wgpu multi-burst", reason);
10811            return None;
10812        }
10813        let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
10814        if !graph_on || self.graph_refused() {
10815            // Same memo as the decode site: this path builds the very
10816            // same graph, so a model it cannot build for must not be
10817            // walked again here either. Missing this guard was worth
10818            // 2.5x on an Adreno — 0.361 tok/s against 0.905 — because
10819            // the burst retried per token what decode had already given
10820            // up on.
10821            return None;
10822        }
10823        let emb = self.embed_single(t_next);
10824        let mut lg = Vec::new();
10825        let mut ids = Vec::new();
10826        match self.try_token_graph_wgpu_steps(
10827            &emb,
10828            position,
10829            &mut lg,
10830            k,
10831            Some(&mut ids),
10832            None,
10833            0,
10834            self.num_layers,
10835        ) {
10836            Some(Ok(_)) => {}
10837            Some(Err(())) => {
10838                // Preserve the backend's post-admission failure through the
10839                // Option-based burst API.  The decode caller consumes this
10840                // flag and clears the sequence instead of falling through
10841                // to a stale CPU recurrent state.
10842                self.graph_failed
10843                    .store(true, std::sync::atomic::Ordering::Relaxed);
10844                return None;
10845            }
10846            None => return None,
10847        }
10848        (ids.len() == k).then_some(ids)
10849    }
10850
10851    /// Multi-step greedy: k whole frames in ONE submit, argmax and re-embed
10852    /// on the device. `ids_out` receives the k winner ids; the hidden/logits
10853    /// outputs are NOT produced in that mode.
10854    fn try_token_graph_wgpu_steps(
10855        &self,
10856        hidden: &[f32],
10857        position: usize,
10858        logits_out: &mut Vec<f32>,
10859        steps: usize,
10860        ids_out: Option<&mut Vec<u32>>,
10861        layers_run: Option<&mut usize>,
10862        from: usize,
10863        upto_excl: usize,
10864    ) -> Option<Result<Vec<f32>, ()>> {
10865        // The bank has already reserved its VRAM. Never build a second
10866        // all-expert arena across bank-owned layers (including bursts).
10867        let upto_excl = match self.mimo_moe.graph_prefix_end() {
10868            Some(end) if end < upto_excl => {
10869                if steps != 1 || layers_run.is_none() || from >= end {
10870                    return None;
10871                }
10872                end
10873            }
10874            _ => upto_excl,
10875        };
10876        // O(1) Nyström decode runs off the sealed state, not the KV cache the
10877        // graph mirrors — never take the graph while o1 is active.
10878        let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
10879        if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
10880            // Softcapped scores have no graph kernel yet — CPU owns them.
10881            // o1 rides the graph only behind CMF_O1_GPU=1 while the port
10882            // proves itself; without it the CPU path owns o1 as before.
10883            return None;
10884        }
10885        // Per-layer KV heads, narrow V, sinks and sliding windows ride
10886        // `GraphAttn::Full::geom` (the ATTEND_X kernels). Anything that
10887        // geometry cannot express declines here, by name — before the
10888        // per-layer gate existed a sliding/sink model ran the graph as
10889        // full-context attention, fluent and wrong. The caller memoizes
10890        // the refusal.
10891        if let Some(reason) = self.wgpu_graph_attn_decline() {
10892            self.note_graph_decline("wgpu token graph", reason);
10893            return None;
10894        }
10895        // Per-layer sealed o1 state for the graph. During prefill the
10896        // state is still Collecting -> views are None -> the graph
10897        // refuses below and the CPU prefill records the q trace and
10898        // seals, exactly as the o1 design requires.
10899        let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
10900            .map(|li| {
10901                if !o1_gpu {
10902                    return None;
10903                }
10904                self.kv_cache.layers[self.phys_layer(li)].o1_views()
10905            })
10906            .collect();
10907        if self.o1_active() && o1_gpu {
10908            // Any o1 layer not sealed (or degenerate exact-only) keeps the
10909            // whole token on the CPU: half-graph forwards would desync.
10910            let want: usize = (from..upto_excl)
10911                .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
10912                .count();
10913            let have = o1_views.iter().filter(|v| v.is_some()).count();
10914            if want == 0 || have != want {
10915                // The silent twin of the gpu-side o1 gates, found the
10916                // same way: a 15x decode drop with an empty log. Views
10917                // stay None until the layer's state SEALS, so `have`
10918                // lagging `want` early in a run is the o1 design working
10919                // — but it must say so, or the next reader spends a
10920                // night proving the kernels innocent.
10921                // On CHANGE, not once: the first decline is the legal
10922                // unsealed prefill, and a once-print buries the state
10923                // that matters — what the count reads AFTER the seal.
10924                use std::sync::atomic::{AtomicUsize, Ordering};
10925                static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
10926                let code = have * 1000 + want;
10927                if LAST.swap(code, Ordering::Relaxed) != code {
10928                    tracing::warn!(
10929                        "o1 graph: {have} of {want} layers sealed — per-op until all seal"
10930                    );
10931                }
10932                return None;
10933            }
10934        }
10935        let nh = self.num_heads;
10936        let (nkv, hd, rd) = self.layer_geom(0);
10937        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
10938        let mut layers = Vec::with_capacity(upto_excl - from);
10939        let mut model = None;
10940        let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
10941        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
10942            if let Some((m, i, kind, rs)) = t
10943                .graph_weight()
10944                .or_else(|| t.graph_weight_descriptor())
10945            {
10946                let name = &m.tensors[i].name;
10947                let prism = if crate::prism::is_inverse_embedding(m, name) {
10948                    crate::gpu::GraphPrismOp::InverseEmbedding
10949                } else if crate::prism::is_forward_weight(m, name) {
10950                    crate::gpu::GraphPrismOp::Forward
10951                } else {
10952                    crate::gpu::GraphPrismOp::None
10953                };
10954                return Some(crate::gpu::GraphW {
10955                    idx: i,
10956                    kind,
10957                    row_scale: rs,
10958                    data: &[],
10959                    prism,
10960                    affine: crate::prism::is_affine_target(m, name),
10961                });
10962            }
10963            // Small unquantized projections (GDN in_proj_a/b) stay f32.
10964            match t.as_f32() {
10965                Some(d) => Some(crate::gpu::GraphW {
10966                    idx: 0,
10967                    kind: 4,
10968                    row_scale: &[],
10969                    data: d,
10970                    prism: crate::gpu::GraphPrismOp::None,
10971                    affine: false,
10972                }),
10973                None => {
10974                    if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
10975                        eprintln!("batch graph: weight has no graph/f32 representation");
10976                    }
10977                    None
10978                }
10979            }
10980        }
10981        for li in from..upto_excl {
10982            let lw = &self.weights.layers[self.phys_layer(li)];
10983            if dbg {
10984                let ak = match &lw.attn {
10985                    AttnKind::Mla(_) => "Mla".into(),
10986                    AttnKind::Full {
10987                        output_gate, bias, ..
10988                    } => format!("Full gate={output_gate} bias={}", bias.is_some()),
10989                    AttnKind::LinearGdn(_) => "LinearGdn".into(),
10990                    AttnKind::Kda(_) => "Kda".into(),
10991                    AttnKind::Linear(_) => "Linear".into(),
10992                    AttnKind::ShortConv(_) => "ShortConv".into(),
10993                    AttnKind::Bounded(_) => "Bounded".into(),
10994                };
10995                let fk = match &lw.ffn {
10996                    FfnKind::Dense(_) => "Dense",
10997                    FfnKind::Moe(_) => "Moe",
10998                    FfnKind::DenseMoe(_) => "DenseMoe",
10999                };
11000                eprintln!("graph L{li}: attn={ak} ffn={fk}");
11001            }
11002            let gffn = match &lw.ffn {
11003                FfnKind::DenseMoe(_) => return None, // dual branch: CPU path
11004                // A tube layer is several matrices, not one — the
11005                // whole-layer graph has no shape for it yet.
11006                FfnKind::Dense(d) if !d.segs.is_empty() => return None,
11007                FfnKind::Dense(d) => {
11008                    // An activation the graph has no kernel arm for keeps
11009                    // the CPU/per-op owner, by name — the dense graph FFN
11010                    // used to compute SiLU for whatever the model asked.
11011                    let Some(act) = d.act.graph_act() else {
11012                        self.note_graph_decline(
11013                            "wgpu token graph",
11014                            "dense FFN activation without a graph kernel",
11015                        );
11016                        return None;
11017                    };
11018                    crate::gpu::GraphFfn::Dense {
11019                        gate: gw(&d.gate_proj)?,
11020                        up: gw(&d.up_proj)?,
11021                        down: gw(&d.down_proj)?,
11022                        act,
11023                    }
11024                }
11025                FfnKind::Moe(m) => {
11026                    // Adaptive τ and expert masks keep the CPU path, where
11027                    // they are implemented. Sigmoid routing with a selection
11028                    // bias (LFM2-MoE / DeepSeek noaux_tc), a routed scale ≠ 1
11029                    // and an UNGATED shared expert (HunYuan hy_v3: ×2.826 on
11030                    // the routed mix, the shared expert at weight 1) are all
11031                    // graphed — before, every such token fell to the per-op
11032                    // path whole (145 submits/token on Hy-MT2-30B-A3B).
11033                    if m.route_tau.is_some() || m.mask.is_some() {
11034                        return None;
11035                    }
11036                    let shared = m.shared.as_ref();
11037                    let has_shared = shared.is_some();
11038                    let shared_gated = matches!(shared, Some((_, Some(_))));
11039                    let sgate = match shared {
11040                        Some((_, Some(sg))) => gw(sg)?,
11041                        // No gate (hy_v3) or no shared expert at all: the
11042                        // router weight stands in so the plumbing stays
11043                        // total; the select kernels pin weight 1 or skip.
11044                        _ => gw(&m.router)?,
11045                    };
11046                    let router = gw(&m.router)?;
11047                    // The resident MoE kernels do not yet carry the
11048                    // descriptor-aware transform through router/shared-gate
11049                    // selection.  Refuse the complete layer instead of
11050                    // scoring with an untransformed Prism plane (the dense
11051                    // path has an explicit FWHT boundary below).
11052                    if router.prism != crate::gpu::GraphPrismOp::None
11053                        || sgate.prism != crate::gpu::GraphPrismOp::None
11054                        || router.affine
11055                        || sgate.affine
11056                    {
11057                        tracing::warn!(
11058                            "resident MoE declined: Prism/affine router or shared gate transform is not implemented"
11059                        );
11060                        return None;
11061                    }
11062                    let inter = m.experts.first()?.gate_proj.rows();
11063                    let mut experts = Vec::with_capacity(m.experts.len() + 1);
11064                    // q4t or q4tp, but not both in one layer — the kernels
11065                    // are picked per layer, not per expert.
11066                    let mut q4tp: Option<bool> = None;
11067                    // The mixed 2-bit profile: q2tp gate/up over a q4tp
11068                    // down. Uniform across the layer, like `q4tp` itself.
11069                    let mut gu_q2: Option<bool> = None;
11070                    for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
11071                        if !matches!(e.act, Act::Silu)
11072                            || e.gate_proj.rows() != inter
11073                            || e.up_proj.rows() != inter
11074                        {
11075                            return None;
11076                        }
11077                        // Expert tensors are packed into one resident buffer
11078                        // and the MoE kernels have no transform slot per
11079                        // expert.  Keep the CPU/per-op owner for Prism or
11080                        // affine experts rather than silently using raw bytes.
11081                        for expert_weight in [&e.gate_proj, &e.up_proj, &e.down_proj] {
11082                            let Some((em, ei, _, _)) = expert_weight
11083                                .graph_weight()
11084                                .or_else(|| expert_weight.graph_weight_descriptor())
11085                            else {
11086                                return None;
11087                            };
11088                            let name = &em.tensors[ei].name;
11089                            if crate::prism::is_forward_weight(em, name)
11090                                || crate::prism::is_inverse_embedding(em, name)
11091                                || crate::prism::is_affine_target(em, name)
11092                            {
11093                                tracing::warn!(
11094                                    "resident MoE declined: expert Prism/affine transform is not implemented"
11095                                );
11096                                return None;
11097                            }
11098                        }
11099                        let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
11100                            Some((mm, gi)) => (
11101                                mm,
11102                                gi,
11103                                e.up_proj.mapped_q4t()?.1,
11104                                e.down_proj.mapped_q4t()?.1,
11105                                false,
11106                                false,
11107                            ),
11108                            None => match e.gate_proj.mapped_q2tp() {
11109                                Some((mm, gi)) => (
11110                                    mm,
11111                                    gi,
11112                                    e.up_proj.mapped_q2tp()?.1,
11113                                    e.down_proj.mapped_q4tp()?.1,
11114                                    true,
11115                                    true,
11116                                ),
11117                                None => {
11118                                    let (mm, gi) = e.gate_proj.mapped_q4tp()?;
11119                                    (
11120                                        mm,
11121                                        gi,
11122                                        e.up_proj.mapped_q4tp()?.1,
11123                                        e.down_proj.mapped_q4tp()?.1,
11124                                        true,
11125                                        false,
11126                                    )
11127                                }
11128                            },
11129                        };
11130                        if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
11131                        {
11132                            // The shared expert rides in the same packed
11133                            // buffer as the routed ones, so a layer that
11134                            // mixes layouts cannot be indexed by one stride.
11135                            // Say so: the symptom is a whole model quietly
11136                            // running its MoE on the CPU.
11137                            tracing::warn!(
11138                                "MoE layer mixes expert layouts (q4tp={is_p}, q2tp gate/up={is_q2})                                  — every expert of a layer, INCLUDING the shared one, must share                                  a layout. The whole-token graph declines this layer."
11139                            );
11140                            return None;
11141                        }
11142                        model.get_or_insert_with(|| mm.clone());
11143                        experts.push((gi, ui, di));
11144                    }
11145                    crate::gpu::GraphFfn::Moe {
11146                        router,
11147                        shared_gate: sgate,
11148                        experts,
11149                        n_exp: m.experts.len(),
11150                        // CMF_TOPK_PROBE: timing probe only — output is WRONG.
11151                        // Fewer experts shrink the MoE arithmetic while the
11152                        // dispatch count stays identical, which is the only
11153                        // clean way to tell a launch-bound decode from a
11154                        // compute-bound one.
11155                        top_k: std::env::var("CMF_TOPK_PROBE")
11156                            .ok()
11157                            .and_then(|v| v.parse::<usize>().ok())
11158                            .filter(|k| *k > 0 && *k <= m.top_k)
11159                            .unwrap_or(m.top_k),
11160                        inter,
11161                        norm_topk: m.norm_topk_prob,
11162                        q4tp: q4tp?,
11163                        gu_q2: gu_q2.unwrap_or(false),
11164                        sigmoid: m.router_sigmoid,
11165                        bias: m.expert_bias.as_deref(),
11166                        has_shared,
11167                        shared_gated,
11168                        route_scale: m.routed_scaling,
11169                    }
11170                }
11171            };
11172            let attn = match &lw.attn {
11173                AttnKind::Full {
11174                    wq,
11175                    wk,
11176                    wv,
11177                    wo,
11178                    q_norm,
11179                    k_norm,
11180                    output_gate,
11181                    softplus_gate,
11182                    bias,
11183                } => {
11184                    if self.attention_heads_per_layer.is_some() {
11185                        return None;
11186                    }
11187                    // A projected output gate rides the graph in one form:
11188                    // Spark-X2.5's head-wise sigmoid (one g_proj row per Q
11189                    // head). Laguna's softplus gate keeps the CPU path.
11190                    let head_gate = match softplus_gate {
11191                        None => None,
11192                        Some((g, true)) if self.proj_gate_sigmoid && !*output_gate => {
11193                            Some(gw(g)?)
11194                        }
11195                        Some(_) => {
11196                            self.note_graph_decline(
11197                                "wgpu token graph",
11198                                "projected softplus / per-element output gate",
11199                            );
11200                            return None;
11201                        }
11202                    };
11203                    let (m, _, _, _) = wq
11204                        .graph_weight()
11205                        .or_else(|| wq.graph_weight_descriptor())?;
11206                    model = Some(m.clone());
11207                    crate::gpu::GraphAttn::Full {
11208                        wq: gw(wq)?,
11209                        wk: gw(wk)?,
11210                        wv: gw(wv)?,
11211                        wo: gw(wo)?,
11212                        q_norm: q_norm.as_deref(),
11213                        k_norm: k_norm.as_deref(),
11214                        late_qk_norm: self.qk_norm_after_rope,
11215                        bias: bias
11216                            .as_ref()
11217                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
11218                        output_gate: *output_gate,
11219                        cpu_k: self.kv_cache.layers[li].k_heads(),
11220                        cpu_v: self.kv_cache.layers[li].v_heads(),
11221                        geom: self.graph_attn_geom(li),
11222                        head_gate,
11223                    }
11224                }
11225                AttnKind::LinearGdn(w) => {
11226                    let cfg = self.gdn_cfg?;
11227                    let (m, _, _, _) = w
11228                        .in_proj_qkv
11229                        .graph_weight()
11230                        .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
11231                    model = Some(m.clone());
11232                    crate::gpu::GraphAttn::Gdn {
11233                        qkv: gw(&w.in_proj_qkv)?,
11234                        z: gw(&w.in_proj_z)?,
11235                        a: gw(&w.in_proj_a)?,
11236                        b: gw(&w.in_proj_b)?,
11237                        out: gw(&w.out_proj)?,
11238                        conv1d: &w.conv1d,
11239                        a_log: &w.a_log,
11240                        dt_bias: &w.dt_bias,
11241                        norm: &w.norm,
11242                        nv: cfg.num_v_heads,
11243                        nk: cfg.num_k_heads,
11244                        dk: cfg.key_head_dim,
11245                        dv: cfg.value_head_dim,
11246                        kk: cfg.conv_kernel,
11247                        cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11248                    }
11249                }
11250                AttnKind::ShortConv(w) => {
11251                    let cfg = self.short_conv_cfg?;
11252                    let (m, _, _, _) = w
11253                        .in_proj
11254                        .graph_weight()
11255                        .or_else(|| w.in_proj.graph_weight_descriptor())?;
11256                    model = Some(m.clone());
11257                    crate::gpu::GraphAttn::ShortConv {
11258                        inp: gw(&w.in_proj)?,
11259                        out: gw(&w.out_proj)?,
11260                        taps: &w.conv,
11261                        kernel: cfg.kernel,
11262                        cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11263                    }
11264                }
11265                _ => return None,
11266            };
11267            layers.push(crate::gpu::GraphLayer {
11268                input_norm: &lw.input_norm,
11269                attn,
11270                post_norm: &lw.post_norm,
11271                ffn: gffn,
11272            });
11273        }
11274        let model = model?;
11275        // Fold final-norm + lm_head into the graph when this call wants logits
11276        // and the lm_head is a graphable (quantized) weight — the graph then
11277        // reads back logits (into logits_out) instead of the hidden, dropping
11278        // the separate CPU/GPU lm_head op + its sync. Never the f32 fallback:
11279        // an unquantized lm_head is vocab·hidden and must not be uploaded.
11280        let lm_gw = if upto_excl == self.num_layers
11281            && self.graph_want_logits
11282            && std::env::var("CMF_GPU_LMHEAD")
11283                .map(|v| v != "0")
11284                .unwrap_or(true)
11285        {
11286            self.weights
11287                .lm_head
11288                .graph_weight()
11289                .or_else(|| self.weights.lm_head.graph_weight_descriptor())
11290                .map(|(m, i, kind, rs)| {
11291                let name = &m.tensors[i].name;
11292                let prism = if crate::prism::is_inverse_embedding(m, name) {
11293                    crate::gpu::GraphPrismOp::InverseEmbedding
11294                } else if crate::prism::is_forward_weight(m, name) {
11295                    crate::gpu::GraphPrismOp::Forward
11296                } else {
11297                    crate::gpu::GraphPrismOp::None
11298                };
11299                (
11300                    crate::gpu::GraphW {
11301                        idx: i,
11302                        kind,
11303                        row_scale: rs,
11304                        data: &[],
11305                        prism,
11306                        affine: crate::prism::is_affine_target(m, name),
11307                    },
11308                    self.weights.lm_head.rows(),
11309                )
11310            })
11311        } else {
11312            None
11313        };
11314        let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
11315        // Multi-step re-embeds the winner on the device.
11316        let emb_gw = if steps > 1 {
11317            self.weights
11318                .embed_tokens
11319                .graph_weight()
11320                .or_else(|| self.weights.embed_tokens.graph_weight_descriptor())
11321                .map(|(m, i, kind, rs)| {
11322                    let name = &m.tensors[i].name;
11323                    let prism = if crate::prism::is_inverse_embedding(m, name) {
11324                        crate::gpu::GraphPrismOp::InverseEmbedding
11325                    } else if crate::prism::is_forward_weight(m, name) {
11326                        crate::gpu::GraphPrismOp::Forward
11327                    } else {
11328                        crate::gpu::GraphPrismOp::None
11329                    };
11330                    (
11331                        crate::gpu::GraphW {
11332                            idx: i,
11333                            kind,
11334                            row_scale: rs,
11335                            data: &[],
11336                            prism,
11337                            affine: crate::prism::is_affine_target(m, name),
11338                        },
11339                        self.weights.embed_tokens.rows(),
11340                        self.embed_multiplier,
11341                    )
11342                })
11343        } else {
11344            None
11345        };
11346
11347        // Loop boundaries: virtual layer indices after which final_norm is
11348        // applied (mid-stack only; the GLOBAL last layer's norm folds into
11349        // lm_head). Span-relative — the executor compares its enumerate
11350        // index. A span ending mid-stack keeps its boundary norm even when
11351        // it is the span's own last layer.
11352        let loop_norm_at: Vec<usize> = if self.loop_final_norm {
11353            (from..upto_excl.min(self.num_layers - 1))
11354                .filter(|&li| (li + 1) % self.physical_layers == 0)
11355                .map(|li| li - from)
11356                .collect()
11357        } else {
11358            Vec::new()
11359        };
11360        let mut h = hidden.to_vec();
11361        // The normal decode path only needs the fused lm-head logits.  A
11362        // CMF_LOGIT_DUMP diagnostic, however, promises a prompt-boundary
11363        // post-stack hidden alongside those logits; request the existing
11364        // second readback only for that explicit probe instead of dumping
11365        // the input copy left in `h` by a folded-head graph.
11366        let dump_hidden = std::env::var_os("CMF_LOGIT_DUMP").is_some();
11367        let outcome = crate::gpu::forward_token_graph(
11368            &model,
11369            self.graph_kv_id,
11370            &layers,
11371            &o1_views,
11372            self.o1_epoch,
11373            &self.inv_freq,
11374            &mut h,
11375            nh,
11376            nkv,
11377            hd,
11378            self.attn_scale,
11379            rd,
11380            self.hidden_size,
11381            self.intermediate_size,
11382            position,
11383            self.kv_cache.max_seq_len,
11384            gemma,
11385            self.rms_eps as f32,
11386            lm,
11387            &self.weights.final_norm,
11388            logits_out,
11389            &loop_norm_at,
11390            steps,
11391            emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
11392            ids_out,
11393            layers_run,
11394            from,
11395            dump_hidden,
11396        );
11397        match outcome {
11398            crate::gpu::TokenGraphOutcome::Completed => Some(Ok(h)),
11399            crate::gpu::TokenGraphOutcome::Failed => Some(Err(())),
11400            crate::gpu::TokenGraphOutcome::Declined => None,
11401        }
11402    }
11403
11404    /// Batched prefill: k contiguous prompt positions through the whole wgpu
11405    /// graph in ONE submit (projections/FFN as GEMMs). `hiddens` is [k·hidden]
11406    /// in/out (embeddings in, layer output out); KV mirror / GDN state advance.
11407    /// false ⇒ unsupported → caller keeps the per-position graph.
11408    /// The b-row Metal graph plan for the whole model: every layer as a
11409    /// GDN run or a full-attention item, all-or-nothing (a layer outside the
11410    /// graph's contract → None, the caller runs plain). Shared by the
11411    /// speculative verify and the batched prefill.
11412    #[cfg(target_os = "macos")]
11413    #[allow(clippy::type_complexity)]
11414    fn metal_rows_plan(
11415        &self,
11416    ) -> Option<(
11417        Vec<MetalRowsItem<'_>>,
11418        std::sync::Arc<cortiq_core::CmfModel>,
11419        Option<crate::gpu_metal::GdnGpuCfg>,
11420    )> {
11421        use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
11422        let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
11423        if !graph_force
11424            || !crate::gpu::enabled_here()
11425            || std::env::var("CMF_GPU_BLOCK")
11426                .map(|v| v == "0")
11427                .unwrap_or(false)
11428            || self.attn_softcap > 0.0
11429            || self.o1_active()
11430            || self.swa.is_some()
11431            || self.global_attn.is_some()
11432            || self.attention_heads_per_layer.is_some()
11433            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
11434            || self.graph_attn_decline_reason().is_some()
11435            || self.attn_v_norm
11436            || self.loop_final_norm
11437        {
11438            return None;
11439        }
11440        let attend_contract = self.head_dim % 4 == 0
11441            && self.head_dim <= 256
11442            && self.rotary_dim >= 2
11443            && self.rotary_dim <= self.head_dim
11444            && (self.rotary_dim / 2) % 32 == 0
11445            && self.num_kv_heads > 0
11446            && self.num_heads % self.num_kv_heads == 0;
11447        if !attend_contract {
11448            return None;
11449        }
11450        let mut plan: Vec<MetalRowsItem> = Vec::new();
11451        let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
11452        for li in 0..self.num_layers {
11453            let lw = &self.weights.layers[self.phys_layer(li)];
11454            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
11455                return None;
11456            }
11457            let ffn = match &lw.ffn {
11458                FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
11459                    let (Some(g), Some(u), Some(dn)) = (
11460                        d.gate_proj.metal_graph_parts(),
11461                        d.up_proj.metal_graph_parts(),
11462                        d.down_proj.metal_graph_parts(),
11463                    ) else {
11464                        return None;
11465                    };
11466                    MetalFfn::Dense {
11467                        gate: g,
11468                        up: u,
11469                        down: dn,
11470                        gelu: false, // SiLU only (the arm above)
11471                    }
11472                }
11473                _ => return None,
11474            };
11475            match &lw.attn {
11476                AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
11477                    let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
11478                        w.in_proj_qkv.metal_graph_parts(),
11479                        w.in_proj_z.metal_graph_parts(),
11480                        w.in_proj_a.f32_parts(),
11481                        w.in_proj_b.f32_parts(),
11482                        w.out_proj.metal_graph_parts(),
11483                    ) else {
11484                        return None;
11485                    };
11486                    if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
11487                        model_ref.get_or_insert_with(|| model.clone());
11488                    }
11489                    let gl = GdnGpuLayer {
11490                        attn_norm: &lw.input_norm,
11491                        post_norm: &lw.post_norm,
11492                        qkv,
11493                        z,
11494                        a,
11495                        b: bb,
11496                        out,
11497                        ffn,
11498                        conv1d: &w.conv1d,
11499                        a_log: &w.a_log,
11500                        dt_bias: &w.dt_bias,
11501                        gnorm: &w.norm,
11502                    };
11503                    match plan.last_mut() {
11504                        Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
11505                        _ => plan.push(MetalRowsItem::Gdn {
11506                            run: vec![gl],
11507                            first: li,
11508                        }),
11509                    }
11510                }
11511                AttnKind::Full {
11512                    wq,
11513                    wk,
11514                    wv,
11515                    wo,
11516                    q_norm,
11517                    k_norm,
11518                    output_gate,
11519                    softplus_gate: None,
11520                    bias: None,
11521                } => {
11522                    let (Some(pq), Some(pk), Some(pv), Some(po)) =
11523                        (
11524                            wq.metal_graph_parts(),
11525                            wk.metal_graph_parts(),
11526                            wv.metal_graph_parts(),
11527                            wo.metal_graph_parts(),
11528                        )
11529                    else {
11530                        return None;
11531                    };
11532                    if let QTensor::Mapped { model, .. } = wq {
11533                        model_ref.get_or_insert_with(|| model.clone());
11534                    }
11535                    let cache = &self.kv_cache.layers[li];
11536                    if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
11537                        return None;
11538                    }
11539                    plan.push(MetalRowsItem::Attn {
11540                        l: AttnGpuLayer {
11541                            attn_norm: &lw.input_norm,
11542                            post_norm: &lw.post_norm,
11543                            wq: pq,
11544                            wk: pk,
11545                            wv: pv,
11546                            wo: po,
11547                            ffn,
11548                        },
11549                        li,
11550                        q_norm: q_norm.as_deref(),
11551                        k_norm: k_norm.as_deref(),
11552                        output_gate: *output_gate,
11553                    });
11554                }
11555                _ => return None,
11556            }
11557        }
11558        let model = model_ref?;
11559        let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
11560            nv: cfg.num_v_heads,
11561            nk: cfg.num_k_heads,
11562            dk: cfg.key_head_dim,
11563            dv: cfg.value_head_dim,
11564            kk: cfg.conv_kernel,
11565            hidden: self.hidden_size,
11566            inter: self.intermediate_size,
11567            c_dim: cfg.conv_dim(),
11568            eps: cfg.rms_eps as f32,
11569            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11570        });
11571        Some((plan, model, gcfg))
11572    }
11573
11574    /// `AttnDeviceParams` for a plan item over the CPU cache as it stands.
11575    #[cfg(target_os = "macos")]
11576    #[allow(clippy::too_many_arguments)]
11577    fn metal_attn_params<'a>(
11578        li: usize,
11579        cache: &'a crate::kv_cache::LayerKvCache,
11580        q_norm: Option<&'a [f32]>,
11581        k_norm: Option<&'a [f32]>,
11582        output_gate: bool,
11583        inv_freq: &'a [f32],
11584        geom: (usize, usize, usize, usize),
11585        pos0: usize,
11586        kv_id: u64,
11587        scale: f32,
11588        eps: f32,
11589        gemma: bool,
11590        late_qk_norm: bool,
11591    ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
11592        let (nh, nkv, hd, rd) = geom;
11593        let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11594        let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11595        let cpu_stored = cpu_k[0].len() / hd;
11596        (
11597            crate::gpu_metal::AttnDeviceParams {
11598                kv_id,
11599                layer: li,
11600                nh,
11601                nkv,
11602                hd,
11603                rd,
11604                position: pos0,
11605                scale,
11606                eps,
11607                gemma,
11608                late_qk_norm,
11609                output_gate,
11610                q_norm,
11611                k_norm,
11612                inv_freq,
11613                cpu_k,
11614                cpu_v,
11615                cpu_stored,
11616                o1: None,
11617                window: None,
11618                head_gate: None,
11619            },
11620            cpu_stored,
11621        )
11622    }
11623
11624    /// Run the rows plan over `hiddens` (b rows at `pos0..`): validate,
11625    /// encode every item, optionally the head, sync. Returns the graph
11626    /// (for the commit / state finish) plus the GDN layer indices and the
11627    /// attention layers with the row count they were encoded against.
11628    #[cfg(target_os = "macos")]
11629    #[allow(clippy::type_complexity)]
11630    fn metal_rows_run(
11631        &mut self,
11632        hiddens: &mut [f32],
11633        pos0: usize,
11634        b: usize,
11635        prefill: bool,
11636        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11637        // Greedy verify: (row length scored, the b argmax ids out) — the
11638        // head's argmax runs on the device and the logits plane is NOT
11639        // read back (`spec.2` stays empty).
11640        mut argmax_out: Option<(usize, &mut Vec<u32>)>,
11641    ) -> MetalRowsRun {
11642        use crate::gpu_metal::{GraphDims, VerifyGraph};
11643        // The previous round's commit may still be replaying into the
11644        // trunk GDN owners on the second queue: this graph reads them
11645        // (zero-copy wraps) and may reallocate them below — collect the
11646        // replay first. Normally already complete (the draft chain ran
11647        // in between); a failed replay is terminal like a failed commit.
11648        if !crate::gpu_metal::wait_replay() {
11649            tracing::error!("Metal rows graph: the pending async replay failed");
11650            return MetalRowsRun::Failed;
11651        }
11652        spec_stamp("v.wait");
11653        // Seed the GDN recurrent records the rows graph reads — GDN layers
11654        // ONLY (the record the CPU path would allocate anyway).  Sizing
11655        // every layer here planted a zero GDN-sized record on a bounded
11656        // anchor of a natively bounded file this path then refused, and
11657        // that record counted as recurrent state on macOS.
11658        let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
11659        if want > 0 {
11660            let phys = self.physical_layers.max(1);
11661            for (li, l) in self.kv_cache.layers.iter_mut().enumerate() {
11662                let is_gdn = self
11663                    .weights
11664                    .layers
11665                    .get(li % phys)
11666                    .is_some_and(|lw| matches!(lw.attn, AttnKind::LinearGdn(_)));
11667                if is_gdn && l.linear_state.len() != want {
11668                    l.linear_state = vec![0f32; want];
11669                }
11670            }
11671        }
11672        let Some((plan, model, gcfg)) = self.metal_rows_plan() else {
11673            return MetalRowsRun::Declined;
11674        };
11675        spec_stamp("v.plan");
11676        let dims = GraphDims {
11677            hidden: self.hidden_size,
11678            eps: self.rms_eps as f32,
11679            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11680        };
11681        let Some(mut graph) = (if prefill {
11682            VerifyGraph::new_prefill(&model, dims, hiddens, b)
11683        } else {
11684            VerifyGraph::new(&model, dims, hiddens, b)
11685        }) else {
11686            return MetalRowsRun::Declined;
11687        };
11688        let geom = (
11689            self.num_heads,
11690            self.num_kv_heads,
11691            self.head_dim,
11692            self.rotary_dim,
11693        );
11694        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11695        let eps = self.rms_eps as f32;
11696        let kv_id = self.graph_kv_id;
11697        let inv_freq = self.inv_freq.clone();
11698        for item in &plan {
11699            let ok = match item {
11700                MetalRowsItem::Gdn { run, .. } => gcfg
11701                    .as_ref()
11702                    .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
11703                    .unwrap_or(false),
11704                MetalRowsItem::Attn {
11705                    l,
11706                    li,
11707                    q_norm,
11708                    k_norm,
11709                    output_gate,
11710                } => {
11711                    let (p, _) = Self::metal_attn_params(
11712                        *li,
11713                        &self.kv_cache.layers[*li],
11714                        *q_norm,
11715                        *k_norm,
11716                        *output_gate,
11717                        &inv_freq,
11718                        geom,
11719                        pos0,
11720                        kv_id,
11721                        self.attn_scale,
11722                        eps,
11723                        gemma,
11724                        self.qk_norm_after_rope,
11725                    );
11726                    graph.attn_ok(l, &p)
11727                }
11728            };
11729            if !ok {
11730                use std::sync::atomic::{AtomicBool, Ordering};
11731                static SAID: AtomicBool = AtomicBool::new(false);
11732                if !SAID.swap(true, Ordering::Relaxed) {
11733                    tracing::warn!("metal rows graph: a layer failed preflight — declining");
11734                }
11735                return MetalRowsRun::Declined;
11736            }
11737        }
11738        let lm = match &spec {
11739            Some((lm, _, _)) => {
11740                if !graph.lm_head_ok(*lm) {
11741                    return MetalRowsRun::Declined;
11742                }
11743                Some(*lm)
11744            }
11745            None => None,
11746        };
11747        let mut gdn_layers = Vec::new();
11748        let mut attn_layers = Vec::new();
11749        for item in &plan {
11750            match item {
11751                MetalRowsItem::Gdn { run, first } => {
11752                    let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
11753                        .iter()
11754                        .map(|l| l.linear_state.as_slice())
11755                        .collect();
11756                    if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
11757                        return MetalRowsRun::Declined;
11758                    }
11759                    gdn_layers.extend(*first..*first + run.len());
11760                }
11761                MetalRowsItem::Attn {
11762                    l,
11763                    li,
11764                    q_norm,
11765                    k_norm,
11766                    output_gate,
11767                } => {
11768                    let (p, cpu_stored) = Self::metal_attn_params(
11769                        *li,
11770                        &self.kv_cache.layers[*li],
11771                        *q_norm,
11772                        *k_norm,
11773                        *output_gate,
11774                        &inv_freq,
11775                        geom,
11776                        pos0,
11777                        kv_id,
11778                        self.attn_scale,
11779                        eps,
11780                        gemma,
11781                        self.qk_norm_after_rope,
11782                    );
11783                    if !graph.encode_attn_b(l, &p) {
11784                        return MetalRowsRun::Declined;
11785                    }
11786                    attn_layers.push((*li, cpu_stored));
11787                }
11788            }
11789        }
11790        if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
11791            if !graph.encode_lm_head_b(final_norm, lm) {
11792                return MetalRowsRun::Declined;
11793            }
11794            // The device argmax is an OPTIMISATION, never a reason to
11795            // decline the round: if it will not encode, drop it and read
11796            // the logits plane back the old way (the head is encoded
11797            // either way, so the rows are there).
11798            if let Some((n, _)) = argmax_out.as_ref() {
11799                if !graph.encode_argmax_b(*n) {
11800                    argmax_out = None;
11801                }
11802            }
11803        }
11804        spec_stamp("v.enc");
11805        if !graph.sync() {
11806            return MetalRowsRun::Failed;
11807        }
11808        spec_stamp("v.gpu");
11809        match (spec, argmax_out) {
11810            (Some(_), Some((_, ids))) => {
11811                ids.resize(b, 0);
11812                if !graph.read_argmax(ids) {
11813                    return MetalRowsRun::Failed;
11814                }
11815                spec_stamp("v.am");
11816            }
11817            (Some((lm, _, logits)), None) => {
11818                logits.resize(b * lm.1, 0.0);
11819                if !graph.read_logits(logits) {
11820                    return MetalRowsRun::Failed;
11821                }
11822                spec_stamp("v.lg");
11823            }
11824            (None, _) => {}
11825        }
11826        if !graph.read_hidden(hiddens) {
11827            return MetalRowsRun::Failed;
11828        }
11829        spec_stamp("v.hid");
11830        MetalRowsRun::Completed(MetalVerifyPending {
11831            graph,
11832            gdn_layers,
11833            attn_layers,
11834        })
11835    }
11836
11837    /// Native-Metal twin of `try_batch_graph_wgpu`: the b rows through the
11838    /// whole model on the `VerifyGraph` (one submit), the head folded in
11839    /// when `spec` asks; `hiddens` come back as the last layer's output
11840    /// rows, `spec.2` as `[b][lm_rows]` logits. The graph is parked in
11841    /// `metal_verify` for `metal_verify_commit`.
11842    #[cfg(target_os = "macos")]
11843    fn try_batch_graph_metal(
11844        &mut self,
11845        hiddens: &mut [f32],
11846        positions: &[usize],
11847        b: usize,
11848        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11849        argmax_out: Option<(usize, &mut Vec<u32>)>,
11850    ) -> crate::gpu::BatchGraphOutcome {
11851        let _t0 = std::time::Instant::now();
11852        if positions.len() != b
11853            || positions.windows(2).any(|w| w[1] != w[0] + 1)
11854            || hiddens.len() != b * self.hidden_size
11855        {
11856            return crate::gpu::BatchGraphOutcome::Declined;
11857        }
11858        let pending = match self.metal_rows_run(hiddens, positions[0], b, false, spec, argmax_out) {
11859            MetalRowsRun::Declined => return crate::gpu::BatchGraphOutcome::Declined,
11860            MetalRowsRun::Failed => return crate::gpu::BatchGraphOutcome::Failed,
11861            MetalRowsRun::Completed(pending) => pending,
11862        };
11863        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
11864            eprintln!(
11865                "metal-verify: {:.1} ms | b={b}",
11866                _t0.elapsed().as_secs_f64() * 1e3
11867            );
11868        }
11869        self.metal_verify = Some(pending);
11870        crate::gpu::BatchGraphOutcome::Completed
11871    }
11872
11873    /// Batched prefill on the Metal rows graph: `ids` (≤ 512) at
11874    /// `start_pos..`, states written in place, K/V rows appended to the
11875    /// CPU caches; optional final norm/head logits are returned in `spec`.
11876    /// Declined means no command buffer was admitted; Failed is terminal.
11877    #[cfg(target_os = "macos")]
11878    fn prefill_rows_metal(
11879        &mut self,
11880        ids: &[u32],
11881        start_pos: usize,
11882        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11883    ) -> MetalPrefillOutcome {
11884        let b = ids.len();
11885        if b == 0 || b > 512 {
11886            return MetalPrefillOutcome::Declined;
11887        }
11888        METAL_PREFILL_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11889        let with_head = spec.is_some();
11890        let hs = self.hidden_size;
11891        let mut hiddens = vec![0f32; b * hs];
11892        for (j, &id) in ids.iter().enumerate() {
11893            let e = self.embed_single(id);
11894            hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
11895        }
11896        let mut pending = match self.metal_rows_run(&mut hiddens, start_pos, b, true, spec, None) {
11897            MetalRowsRun::Declined => return MetalPrefillOutcome::Declined,
11898            MetalRowsRun::Failed => {
11899                METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11900                return MetalPrefillOutcome::Failed;
11901            }
11902            MetalRowsRun::Completed(pending) => pending,
11903        };
11904        // states are final: copy them to the owners
11905        let idxs = pending.gdn_layers.clone();
11906        let mut outs: Vec<&mut [f32]> = self
11907            .kv_cache
11908            .layers
11909            .iter_mut()
11910            .enumerate()
11911            .filter(|(i, _)| idxs.binary_search(i).is_ok())
11912            .map(|(_, l)| l.linear_state.as_mut_slice())
11913            .collect();
11914        if !pending.graph.finish_states(&mut outs) {
11915            METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11916            return MetalPrefillOutcome::Failed;
11917        }
11918        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
11919        // Read every layer before mutating any CPU cache.  A missing mirror
11920        // row is a terminal graph failure, not a reason to append a partial
11921        // prefix and replay the remainder serially.
11922        let mut rows = Vec::with_capacity(pending.attn_layers.len());
11923        for (li, cpu_stored) in &pending.attn_layers {
11924            let mut kbuf = vec![0f32; b * nkv * hd];
11925            let mut vbuf = vec![0f32; b * nkv * hd];
11926            if !crate::gpu_metal::kv_mirror_read_rows(
11927                self.graph_kv_id,
11928                *li,
11929                nkv,
11930                hd,
11931                *cpu_stored,
11932                b,
11933                &mut kbuf,
11934                &mut vbuf,
11935            ) {
11936                METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11937                return MetalPrefillOutcome::Failed;
11938            }
11939            rows.push((*li, *cpu_stored, kbuf, vbuf));
11940        }
11941        for (li, cpu_stored, kbuf, vbuf) in rows {
11942            let cache = &mut self.kv_cache.layers[li];
11943            for r in 0..b {
11944                cache.append(
11945                    &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
11946                    &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
11947                    &[],
11948                );
11949            }
11950            crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + b);
11951        }
11952        METAL_PREFILL_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11953        if with_head {
11954            METAL_PREFILL_HEAD_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11955        }
11956        MetalPrefillOutcome::Completed(hiddens)
11957    }
11958
11959    #[cfg(target_os = "macos")]
11960    fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> MetalPrefillOutcome {
11961        self.prefill_rows_metal(ids, start_pos, None)
11962    }
11963
11964    /// Exact teacher-forced NLL through the ordinary Metal rows graph.  This
11965    /// is intentionally separate from the serial TokenGraph scorer: every
11966    /// chunk owns a real b-row graph/head completion and the recurrent/KV
11967    /// handoff is committed before the next chunk begins.
11968    #[cfg(target_os = "macos")]
11969    fn nll_batch_metal(&mut self, ids: &[u32], start: usize) -> MetalBatchNllOutcome {
11970        if ids.len() < 2 || self.o1_active() || self.head_clusters.is_some() {
11971            return MetalBatchNllOutcome::Declined;
11972        }
11973        let Some(lm) = self.weights.lm_head.metal_graph_parts() else {
11974            return MetalBatchNllOutcome::Declined;
11975        };
11976        let chunk = std::env::var("CMF_METAL_PREFILL_CHUNK")
11977            .ok()
11978            .and_then(|v| v.parse::<usize>().ok())
11979            .filter(|&v| (1..=512).contains(&v))
11980            .unwrap_or(32);
11981        let final_norm = self.weights.final_norm.clone();
11982        let mut nll = 0.0f64;
11983        let mut count = 0usize;
11984        let mut pos = 0usize;
11985        let mut completed = 0usize;
11986        while pos < ids.len() {
11987            let end = (pos + chunk).min(ids.len());
11988            let mut logits = Vec::new();
11989            let outcome = self.prefill_rows_metal(
11990                &ids[pos..end],
11991                pos,
11992                Some((lm, &final_norm, &mut logits)),
11993            );
11994            match outcome {
11995                MetalPrefillOutcome::Declined => {
11996                    return if completed == 0 {
11997                        MetalBatchNllOutcome::Declined
11998                    } else {
11999                        MetalBatchNllOutcome::Failed(format!(
12000                            "ordinary Metal NLL batch declined after {completed} chunks"
12001                        ))
12002                    };
12003                }
12004                MetalPrefillOutcome::Failed => {
12005                    return MetalBatchNllOutcome::Failed(
12006                        "ordinary Metal NLL batch failed after admission".to_string(),
12007                    );
12008                }
12009                MetalPrefillOutcome::Completed(_) => {}
12010            }
12011            completed += 1;
12012            let vocab = self.vocab_size.min(lm.1);
12013            if logits.len() != (end - pos) * lm.1 || vocab == 0 {
12014                return MetalBatchNllOutcome::Failed(
12015                    "ordinary Metal NLL head returned an invalid shape".to_string(),
12016                );
12017            }
12018            for row in 0..(end - pos) {
12019                let absolute = pos + row;
12020                if absolute < start || absolute + 1 >= ids.len() {
12021                    continue;
12022                }
12023                let lg = &mut logits[row * lm.1..row * lm.1 + vocab];
12024                if let Some(mu) = self.logit_multiplier {
12025                    for v in lg.iter_mut() {
12026                        *v *= mu;
12027                    }
12028                }
12029                if let Some(c) = self.final_softcap {
12030                    for v in lg.iter_mut() {
12031                        *v = c * (*v / c).tanh();
12032                    }
12033                }
12034                let target = ids[absolute + 1] as usize;
12035                if target >= vocab {
12036                    return MetalBatchNllOutcome::Failed(format!(
12037                        "target token {target} exceeds Metal head rows {vocab}"
12038                    ));
12039                }
12040                let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
12041                let lse: f64 = lg
12042                    .iter()
12043                    .map(|&v| ((v - max) as f64).exp())
12044                    .sum::<f64>()
12045                    .ln()
12046                    + max as f64;
12047                nll += lse - lg[target] as f64;
12048                count += 1;
12049            }
12050            pos = end;
12051        }
12052        MetalBatchNllOutcome::Completed(nll, count)
12053    }
12054
12055    /// Commit a Metal verify round: replay the GDN recurrences over the
12056    /// `a + 1` accepted positions into the CPU states, append the accepted
12057    /// K/V rows from the mirrors to the CPU caches, re-point the mirrors.
12058    #[cfg(target_os = "macos")]
12059    fn metal_verify_commit(&mut self, a: usize) -> bool {
12060        let Some(mut pending) = self.metal_verify.take() else {
12061            return false;
12062        };
12063        let n = a + 1;
12064        // encode order == ascending layer order (the plan walks 0..layers)
12065        let idxs = pending.gdn_layers.clone();
12066        let mut outs: Vec<&mut [f32]> = self
12067            .kv_cache
12068            .layers
12069            .iter_mut()
12070            .enumerate()
12071            .filter(|(i, _)| idxs.binary_search(i).is_ok())
12072            .map(|(_, l)| l.linear_state.as_mut_slice())
12073            .collect();
12074        if !pending.graph.commit(n, &mut outs) {
12075            return false;
12076        }
12077        spec_stamp("c.replay");
12078        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12079        // Read every layer before mutating any CPU cache.  Missing rows are
12080        // terminal after the replay has executed; never append a partial KV
12081        // prefix and continue on a serial path.
12082        let mut rows = Vec::with_capacity(pending.attn_layers.len());
12083        for (li, cpu_stored) in &pending.attn_layers {
12084            let mut kbuf = vec![0f32; n * nkv * hd];
12085            let mut vbuf = vec![0f32; n * nkv * hd];
12086            if !crate::gpu_metal::kv_mirror_read_rows(
12087                self.graph_kv_id,
12088                *li,
12089                nkv,
12090                hd,
12091                *cpu_stored,
12092                n,
12093                &mut kbuf,
12094                &mut vbuf,
12095            ) {
12096                return false;
12097            }
12098            rows.push((*li, *cpu_stored, kbuf, vbuf));
12099        }
12100        for (li, cpu_stored, kbuf, vbuf) in rows {
12101            let cache = &mut self.kv_cache.layers[li];
12102            for r in 0..n {
12103                cache.append(
12104                    &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12105                    &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12106                    &[],
12107                );
12108            }
12109            crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + n);
12110        }
12111        spec_stamp("c.kv");
12112        true
12113    }
12114
12115    /// The round's warm-ups as ONE b-row graph run over the MTP block on
12116    /// Metal: `pairs` = (trunk hidden, next token) at consecutive positions
12117    /// from `first_pos`; the block's input projection is folded in. This
12118    /// half encodes and SUBMITS (no wait); `mtp_warm_batch_finish` waits
12119    /// and pulls the appended K/V rows into the CPU MTP cache. None = the
12120    /// graph declined (nothing submitted, nothing appended).
12121    #[cfg(target_os = "macos")]
12122    fn mtp_warm_batch_submit(
12123        &mut self,
12124        m: &mut MtpModule,
12125        pairs: &[(&[f32], u32)],
12126        first_pos: usize,
12127    ) -> Option<MetalWarmPending> {
12128        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
12129        let b = pairs.len();
12130        if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
12131            return None;
12132        }
12133        let AttnKind::Full {
12134            wq,
12135            wk,
12136            wv,
12137            wo,
12138            q_norm,
12139            k_norm,
12140            output_gate,
12141            softplus_gate: None,
12142            bias: None,
12143        } = &m.layer.attn
12144        else {
12145            return None;
12146        };
12147        let FfnKind::Dense(d) = &m.layer.ffn else {
12148            return None;
12149        };
12150        if !d.segs.is_empty() {
12151            return None;
12152        }
12153        let (Some(pq), Some(pk), Some(pv), Some(po)) =
12154            (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12155        else {
12156            return None;
12157        };
12158        let (Some(g), Some(u), Some(dn)) = (
12159            d.gate_proj.q1_parts(),
12160            d.up_proj.q1_parts(),
12161            d.down_proj.q1_parts(),
12162        ) else {
12163            return None;
12164        };
12165        let Some(eh) = m.eh_proj.q1_parts() else {
12166            return None;
12167        };
12168        let QTensor::Mapped { model, .. } = wq else {
12169            return None;
12170        };
12171        let model = model.clone();
12172        let hs = self.hidden_size;
12173        // [enorm(embed(tok)); hnorm(hidden)] rows
12174        let mut cat = vec![0f32; b * 2 * hs];
12175        for (j, (h, tok)) in pairs.iter().enumerate() {
12176            let e = self.embed_single(*tok);
12177            let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
12178            inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
12179            inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
12180        }
12181        let dims = GraphDims {
12182            hidden: hs,
12183            eps: self.rms_eps as f32,
12184            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12185        };
12186        spec_stamp("w.cat");
12187        let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
12188            return None;
12189        };
12190        spec_stamp("w.new");
12191        let l = AttnGpuLayer {
12192            attn_norm: &m.layer.input_norm,
12193            post_norm: &m.layer.post_norm,
12194            wq: pq,
12195            wk: pk,
12196            wv: pv,
12197            wo: po,
12198            ffn: MetalFfn::Dense {
12199                gate: g,
12200                up: u,
12201                down: dn,
12202                gelu: d.act == Act::Gelu,
12203            },
12204        };
12205        let (nh, nkv, hd, rd) = (
12206            self.num_heads,
12207            self.num_kv_heads,
12208            self.head_dim,
12209            self.rotary_dim,
12210        );
12211        let inv_freq = self.inv_freq.clone();
12212        let cpu_stored;
12213        {
12214            let cache = &m.kv;
12215            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12216            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12217            cpu_stored = cpu_k[0].len() / hd;
12218            // The cache may LAG the position (rows nobody warmed): the
12219            // pairs land at cpu_stored.. with their true RoPE positions
12220            // first_pos.., exactly what the one-by-one warm does. A cache
12221            // AHEAD of the position is a real inconsistency.
12222            if cpu_stored > first_pos {
12223                spec_stamp("w.decl");
12224                return None;
12225            }
12226            let p = AttnDeviceParams {
12227                kv_id: self.mtp_kv_id(),
12228                layer: Self::MTP_LAYER_BASE,
12229                nh,
12230                nkv,
12231                hd,
12232                rd,
12233                position: first_pos,
12234                scale: self.attn_scale,
12235                eps: self.rms_eps as f32,
12236                gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12237                late_qk_norm: self.qk_norm_after_rope,
12238                output_gate: *output_gate,
12239                q_norm: q_norm.as_deref(),
12240                k_norm: k_norm.as_deref(),
12241                inv_freq: &inv_freq,
12242                cpu_k,
12243                cpu_v,
12244                cpu_stored,
12245                o1: None,
12246                window: None,
12247                head_gate: None,
12248            };
12249            if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
12250                return None;
12251            }
12252        }
12253        spec_stamp("w.enc");
12254        if !graph.submit() {
12255            return None;
12256        }
12257        spec_stamp("w.sub");
12258        Some(MetalWarmPending {
12259            graph,
12260            cpu_stored,
12261            b,
12262        })
12263    }
12264
12265    /// Submit and finish in one call (the prefill's MTP warm-up, where
12266    /// nothing runs in between).
12267    #[cfg(target_os = "macos")]
12268    fn mtp_warm_batch_metal(
12269        &mut self,
12270        m: &mut MtpModule,
12271        pairs: &[(&[f32], u32)],
12272        first_pos: usize,
12273    ) -> bool {
12274        match self.mtp_warm_batch_submit(m, pairs, first_pos) {
12275            Some(p) => self.mtp_warm_batch_finish(m, p),
12276            None => false,
12277        }
12278    }
12279
12280    /// Second half of the batched warm-up: wait for the submitted graph,
12281    /// pull its b appended K/V rows into the CPU MTP cache, re-point the
12282    /// mirror. False = the command buffer failed or the rows are missing
12283    /// (nothing appended; the caller falls back to the one-by-one warm).
12284    #[cfg(target_os = "macos")]
12285    fn mtp_warm_batch_finish(&mut self, m: &mut MtpModule, pending: MetalWarmPending) -> bool {
12286        let MetalWarmPending {
12287            mut graph,
12288            cpu_stored,
12289            b,
12290        } = pending;
12291        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12292        if !graph.sync() {
12293            return false;
12294        }
12295        spec_stamp("w.gpu");
12296        let mut kbuf = vec![0f32; b * nkv * hd];
12297        let mut vbuf = vec![0f32; b * nkv * hd];
12298        if !crate::gpu_metal::kv_mirror_read_rows(
12299            self.mtp_kv_id(),
12300            Self::MTP_LAYER_BASE,
12301            nkv,
12302            hd,
12303            cpu_stored,
12304            b,
12305            &mut kbuf,
12306            &mut vbuf,
12307        ) {
12308            return false;
12309        }
12310        for r in 0..b {
12311            m.kv.append(
12312                &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12313                &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12314                &[],
12315            );
12316        }
12317        crate::gpu_metal::kv_mirror_set_stored(
12318            self.mtp_kv_id(),
12319            Self::MTP_LAYER_BASE,
12320            cpu_stored + b,
12321        );
12322        spec_stamp("w.kv");
12323        true
12324    }
12325
12326    /// A committed token id from the high table (Cyrillic, CJK and the
12327    /// like sit above 131072 in Qwen's vocabulary; Latin subwords past
12328    /// the 65536 cut are rare enough to lose as rejected drafts) switches
12329    /// the draft to the full head for the next 16 tokens; other ids count
12330    /// down. On an M4 the full 660 MB head costs 5.5 ms a draft step
12331    /// against 1.4 for the shortlist, so the streak is kept short.
12332    pub(crate) fn note_draft_id(&mut self, id: u32) {
12333        let cut = Self::draft_vocab_rows(usize::MAX).max(131_072);
12334        if (id as usize) >= cut {
12335            self.draft_full_streak = 16;
12336        } else {
12337            self.draft_full_streak = self.draft_full_streak.saturating_sub(1);
12338        }
12339    }
12340
12341    /// The draft head's rows for the next step: the shortlist, or the full
12342    /// head while `draft_full_streak` runs.
12343    fn draft_head_rows(&self, head_rows: usize) -> usize {
12344        if self.draft_full_streak > 0 {
12345            head_rows
12346        } else {
12347            Self::draft_vocab_rows(head_rows)
12348        }
12349    }
12350
12351    /// Draft-head shortlist size: `CMF_DRAFT_VOCAB` rows (default 65536,
12352    /// capped at the head; 0 = full head).
12353    fn draft_vocab_rows(head_rows: usize) -> usize {
12354        static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
12355        let n = *N.get_or_init(|| {
12356            std::env::var("CMF_DRAFT_VOCAB")
12357                .ok()
12358                .and_then(|v| v.parse().ok())
12359                .unwrap_or(65536)
12360        });
12361        if n == 0 { head_rows } else { n.min(head_rows) }
12362    }
12363
12364    /// One MTP block step on the native Metal token graph: block input on
12365    /// the host, the attention layer + FFN device-resident over the MTP
12366    /// mirror, the head folded in when `want_logits`. The appended K/V row
12367    /// is pulled into the CPU MTP cache (owner of record) after the sync.
12368    #[cfg(target_os = "macos")]
12369    fn mtp_step_metal(
12370        &mut self,
12371        m: &mut MtpModule,
12372        hidden: &[f32],
12373        next_token: u32,
12374        position: usize,
12375        want_logits: bool,
12376    ) -> Option<(Vec<f32>, Vec<f32>)> {
12377        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12378        if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12379            || !crate::gpu::q1_force()
12380            || !crate::gpu::enabled_here()
12381            || self.attn_softcap > 0.0
12382            || self.attention_heads_per_layer.is_some()
12383            || m.kv.mode != crate::kv_cache::KvMode::F32
12384            || m.kv.o1.is_some()
12385        {
12386            return None;
12387        }
12388        let AttnKind::Full {
12389            wq,
12390            wk,
12391            wv,
12392            wo,
12393            q_norm,
12394            k_norm,
12395            output_gate,
12396            softplus_gate: None,
12397            bias: None,
12398        } = &m.layer.attn
12399        else {
12400            return None;
12401        };
12402        let FfnKind::Dense(d) = &m.layer.ffn else {
12403            return None;
12404        };
12405        if d.act != Act::Silu || !d.segs.is_empty() {
12406            return None;
12407        }
12408        let (pq, pk, pv, po) = (
12409            wq.q1_parts()?,
12410            wk.q1_parts()?,
12411            wv.q1_parts()?,
12412            wo.q1_parts()?,
12413        );
12414        let (g, u, dn) = (
12415            d.gate_proj.q1_parts()?,
12416            d.up_proj.q1_parts()?,
12417            d.down_proj.q1_parts()?,
12418        );
12419        let QTensor::Mapped { model, .. } = wq else {
12420            return None;
12421        };
12422        let model = model.clone();
12423        let lm = if want_logits {
12424            Some(self.weights.lm_head.q1_parts()?)
12425        } else {
12426            None
12427        };
12428        let dims = GraphDims {
12429            hidden: self.hidden_size,
12430            eps: self.rms_eps as f32,
12431            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12432        };
12433        // The block input `eh_proj · [enorm(e); hnorm(h)]` rides in the
12434        // graph (one submit a step); the host per-op matvec if it cannot.
12435        let hs = self.hidden_size;
12436        let mut x = vec![0f32; hs];
12437        let mut graph = TokenGraph::new(&model, dims, &x)?;
12438        let mut folded = false;
12439        if let Some(eh) = m.eh_proj.q1_parts() {
12440            let e = self.embed_single(next_token);
12441            let mut cat = vec![0.0f32; 2 * hs];
12442            let (cat_e, cat_h) = cat.split_at_mut(hs);
12443            inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
12444            inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
12445            folded = graph.encode_input_proj(eh, &cat);
12446        }
12447        if !folded {
12448            x = self.mtp_block_input(m, hidden, next_token);
12449            graph = TokenGraph::new(&model, dims, &x)?;
12450        }
12451        spec_stamp("d.in");
12452        let l = AttnGpuLayer {
12453            attn_norm: &m.layer.input_norm,
12454            post_norm: &m.layer.post_norm,
12455            wq: pq,
12456            wk: pk,
12457            wv: pv,
12458            wo: po,
12459            ffn: MetalFfn::Dense {
12460                gate: g,
12461                up: u,
12462                down: dn,
12463                gelu: d.act == Act::Gelu,
12464            },
12465        };
12466        let (nh, nkv, hd, rd) = (
12467            self.num_heads,
12468            self.num_kv_heads,
12469            self.head_dim,
12470            self.rotary_dim,
12471        );
12472        let inv_freq = self.inv_freq.clone();
12473        {
12474            let cache = &m.kv;
12475            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12476            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12477            let cpu_stored = cpu_k[0].len() / hd;
12478            let p = AttnDeviceParams {
12479                kv_id: self.mtp_kv_id(),
12480                layer: Self::MTP_LAYER_BASE,
12481                nh,
12482                nkv,
12483                hd,
12484                rd,
12485                position,
12486                scale: self.attn_scale,
12487                eps: self.rms_eps as f32,
12488                gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12489                late_qk_norm: self.qk_norm_after_rope,
12490                output_gate: *output_gate,
12491                q_norm: q_norm.as_deref(),
12492                k_norm: k_norm.as_deref(),
12493                inv_freq: &inv_freq,
12494                cpu_k,
12495                cpu_v,
12496                cpu_stored,
12497                o1: None,
12498                window: None,
12499                head_gate: None,
12500            };
12501            if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12502                return None;
12503            }
12504        }
12505        // The draft's head over a vocabulary SHORTLIST (the first
12506        // CMF_DRAFT_VOCAB rows — BPE ids run roughly by merge rank, so the
12507        // low ids carry the mass): the verify keeps the full head, so a true
12508        // token past the cut is only a rejected draft, never a wrong token.
12509        // 662 MB a step on Qwen3.8 becomes 170 MB at 65536.
12510        let draft_rows = if let Some(lm) = lm {
12511            self.draft_head_rows(lm.1)
12512        } else {
12513            0
12514        };
12515        if let Some(lm) = lm {
12516            if !graph.lm_head_ok(lm) {
12517                return None;
12518            }
12519            if draft_rows < lm.1 {
12520                if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12521                    return None;
12522                }
12523            } else {
12524                graph.encode_lm_head(&m.final_norm, lm);
12525            }
12526        }
12527        spec_stamp("d.enc");
12528        if graph.sync_checked().is_err() {
12529            return None;
12530        }
12531        spec_stamp("d.gpu");
12532        let mut logits = Vec::new();
12533        if let Some(lm) = lm {
12534            let n_read = draft_rows.min(lm.1).min(self.vocab_size);
12535            logits = attention::take_buf(n_read);
12536            graph.read_logits(&mut logits);
12537            // ids past the shortlist: never drafted (−∞ in every chain)
12538            logits.resize(self.vocab_size, f32::NEG_INFINITY);
12539        }
12540        graph.finish(&mut x);
12541        let mut krow = attention::take_buf(nkv * hd);
12542        let mut vrow = attention::take_buf(nkv * hd);
12543        if crate::gpu_metal::kv_mirror_read_last(
12544            self.mtp_kv_id(),
12545            Self::MTP_LAYER_BASE,
12546            nkv,
12547            hd,
12548            &mut krow,
12549            &mut vrow,
12550        ) {
12551            m.kv.append(&krow, &vrow, &[]);
12552        }
12553        attention::recycle_buf(&mut krow);
12554        attention::recycle_buf(&mut vrow);
12555        spec_stamp("d.rd");
12556        Some((logits, x))
12557    }
12558
12559    /// `CMF_MTP_CHAIN=0` keeps the per-step draft (one submit and one
12560    /// host round trip per MTP step); the default drafts the whole chain
12561    /// in one command buffer when the round is plain greedy.
12562    ///
12563    /// Measured on an M4 (24 GB), Qwen3.8-27B q4tp, P3 at 160 tokens,
12564    /// k=7, six runs per arm alternating inside one lock window — the
12565    /// round's draft phase (median over the 34 rounds of a run) is
12566    /// 34.5 ms per round old against 30.1 new, i.e. 4.93 → 4.31 ms per
12567    /// draft step. That is the whole prize: the 7 submits cost ~0.6 ms
12568    /// each in host and submit latency and nothing else changes —
12569    /// acceptance (3.41 of 7) and tokens per round (4.41) are identical,
12570    /// and the round is 289 → 285 ms, decode 13.8 → 14.0 tok/s.
12571    fn mtp_chain_on() -> bool {
12572        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12573        *ON.get_or_init(|| std::env::var("CMF_MTP_CHAIN").as_deref() != Ok("0"))
12574    }
12575
12576    /// The round's k greedy drafts as ONE command buffer on Metal: the MTP
12577    /// block k times back to back, each step's token embedding gathered
12578    /// on the device from the argmax the step before it wrote, the head
12579    /// over the round's shortlist (or the full head during a full-head
12580    /// streak — decided once, before the chain, exactly as the per-step
12581    /// path decides it per step, since `draft_full_streak` only moves on
12582    /// a commit). One wait, then the k ids and the k appended K/V rows
12583    /// come back; the CPU MTP cache ends where k `mtp_step_metal` calls
12584    /// would have left it. `Err(false)` = declined before anything was
12585    /// committed (the per-step path takes the round); `Err(true)` = the
12586    /// command buffer failed after commit.
12587    #[cfg(target_os = "macos")]
12588    fn mtp_draft_chain_metal(
12589        &mut self,
12590        m: &mut MtpModule,
12591        hidden: &[f32],
12592        t_next: u32,
12593        position: usize,
12594        k: usize,
12595    ) -> Result<Vec<u32>, bool> {
12596        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12597        if k == 0
12598            || k > 64
12599            || !Self::mtp_chain_on()
12600            || std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12601            || !crate::gpu::q1_force()
12602            || !crate::gpu::enabled_here()
12603            || self.attn_softcap > 0.0
12604            || self.attention_heads_per_layer.is_some()
12605            || m.kv.mode != crate::kv_cache::KvMode::F32
12606            || m.kv.o1.is_some()
12607            // the chain gathers embeddings itself: only the plain table
12608            || self.dsv4.is_some()
12609            || self.dsv41.is_some()
12610            || self.qwen4_exp.is_some()
12611            || self.g3n.is_some()
12612        {
12613            return Err(false);
12614        }
12615        let AttnKind::Full {
12616            wq,
12617            wk,
12618            wv,
12619            wo,
12620            q_norm,
12621            k_norm,
12622            output_gate,
12623            softplus_gate: None,
12624            bias: None,
12625        } = &m.layer.attn
12626        else {
12627            return Err(false);
12628        };
12629        let FfnKind::Dense(d) = &m.layer.ffn else {
12630            return Err(false);
12631        };
12632        if d.act != Act::Silu || !d.segs.is_empty() {
12633            return Err(false);
12634        }
12635        let (Some(pq), Some(pk), Some(pv), Some(po)) =
12636            (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12637        else {
12638            return Err(false);
12639        };
12640        let (Some(g), Some(u), Some(dn)) = (
12641            d.gate_proj.q1_parts(),
12642            d.up_proj.q1_parts(),
12643            d.down_proj.q1_parts(),
12644        ) else {
12645            return Err(false);
12646        };
12647        let (Some(eh), Some(lm)) = (m.eh_proj.q1_parts(), self.weights.lm_head.q1_parts()) else {
12648            return Err(false);
12649        };
12650        let QTensor::Mapped { model, .. } = wq else {
12651            return Err(false);
12652        };
12653        let model = model.clone();
12654        // the embedding table: a q4tp tensor of the SAME blob, no Prism
12655        // inverse-embedding post-pass
12656        let QTensor::Mapped {
12657            model: em,
12658            idx: eidx,
12659            dtype: cortiq_core::TensorDtype::Q4TiledP,
12660            ..
12661        } = &self.weights.embed_tokens
12662        else {
12663            return Err(false);
12664        };
12665        if !std::sync::Arc::ptr_eq(em, &model)
12666            || crate::prism::is_inverse_embedding(&model, &model.tensors[*eidx].name)
12667        {
12668            return Err(false);
12669        }
12670        let embed = (
12671            *eidx,
12672            self.weights.embed_tokens.rows(),
12673            self.weights.embed_tokens.cols(),
12674        );
12675        if embed.2 != self.hidden_size || hidden.len() != self.hidden_size {
12676            return Err(false);
12677        }
12678        let dims = GraphDims {
12679            hidden: self.hidden_size,
12680            eps: self.rms_eps as f32,
12681            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12682        };
12683        let Some(mut graph) = TokenGraph::new(&model, dims, hidden) else {
12684            return Err(false);
12685        };
12686        if !graph.chain_embed_ok(embed) || !graph.lm_head_ok(lm) {
12687            return Err(false);
12688        }
12689        let l = AttnGpuLayer {
12690            attn_norm: &m.layer.input_norm,
12691            post_norm: &m.layer.post_norm,
12692            wq: pq,
12693            wk: pk,
12694            wv: pv,
12695            wo: po,
12696            ffn: MetalFfn::Dense {
12697                gate: g,
12698                up: u,
12699                down: dn,
12700                gelu: d.act == Act::Gelu,
12701            },
12702        };
12703        let (nh, nkv, hd, rd) = (
12704            self.num_heads,
12705            self.num_kv_heads,
12706            self.head_dim,
12707            self.rotary_dim,
12708        );
12709        let inv_freq = self.inv_freq.clone();
12710        let draft_rows = self.draft_head_rows(lm.1);
12711        let n_arg = draft_rows.min(lm.1).min(self.vocab_size);
12712        if n_arg == 0 {
12713            return Err(false);
12714        }
12715        // `CMF_MTP_CHAIN_SPLIT=1` commits each step as it is encoded, so
12716        // the GPU starts on step 0 while the host is still encoding step
12717        // 1 — a probe for whether the host encode is on the critical
12718        // path. It is not: three runs each, draft 30.0 ms per round split
12719        // against 30.1 whole, and the whole chain's host encode measures
12720        // 0.3 ms against a 29.7 ms wait. Kept as a probe, off by default.
12721        let split = std::env::var("CMF_MTP_CHAIN_SPLIT").as_deref() == Ok("1");
12722        let t_chain = std::time::Instant::now();
12723        graph.chain_ids_init(t_next, k);
12724        let cpu_stored;
12725        {
12726            let cache = &m.kv;
12727            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12728            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12729            cpu_stored = cpu_k[0].len() / hd;
12730            for j in 0..k {
12731                if !graph.encode_chain_input(
12732                    embed,
12733                    j as u32,
12734                    &m.enorm,
12735                    &m.hnorm,
12736                    self.embed_multiplier,
12737                    eh,
12738                ) {
12739                    return Err(false);
12740                }
12741                // step j's mirror row: the mirror is re-pointed at the CPU
12742                // rows before step 0 and advances by one per step; its
12743                // resync (never taken past step 0) reads the CPU rows
12744                let p = AttnDeviceParams {
12745                    kv_id: self.mtp_kv_id(),
12746                    layer: Self::MTP_LAYER_BASE,
12747                    nh,
12748                    nkv,
12749                    hd,
12750                    rd,
12751                    position: position + j,
12752                    scale: self.attn_scale,
12753                    eps: self.rms_eps as f32,
12754                    gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12755                    late_qk_norm: self.qk_norm_after_rope,
12756                    output_gate: *output_gate,
12757                    q_norm: q_norm.as_deref(),
12758                    k_norm: k_norm.as_deref(),
12759                    inv_freq: &inv_freq,
12760                    cpu_k: cpu_k.clone(),
12761                    cpu_v: cpu_v.clone(),
12762                    cpu_stored: cpu_stored + j,
12763                    o1: None,
12764                    window: None,
12765                    head_gate: None,
12766                };
12767                if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12768                    return Err(false);
12769                }
12770                if draft_rows < lm.1 {
12771                    if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12772                        return Err(false);
12773                    }
12774                } else {
12775                    graph.encode_lm_head(&m.final_norm, lm);
12776                }
12777                if !graph.encode_argmax(n_arg, j as u32 + 1) {
12778                    return Err(false);
12779                }
12780                if split {
12781                    // CMF_MTP_CHAIN_SPLIT=1: commit every step so the GPU
12782                    // starts on step 0 while the host encodes the rest
12783                    graph.commit();
12784                }
12785            }
12786        }
12787        let t_enc = t_chain.elapsed();
12788        if graph.sync_checked().is_err() {
12789            return Err(true);
12790        }
12791        if std::env::var_os("CMF_GRAPH_SPEC_TIME").is_some() {
12792            eprintln!(
12793                "mtp-chain: encode {:.1} ms | wait {:.1} ms (k={k}, head rows {draft_rows}{})",
12794                t_enc.as_secs_f64() * 1e3,
12795                (t_chain.elapsed() - t_enc).as_secs_f64() * 1e3,
12796                if split { ", split" } else { "" }
12797            );
12798        }
12799        let mut ids = vec![0u32; k];
12800        if !graph.chain_ids_read(&mut ids) {
12801            return Err(true);
12802        }
12803        let mut kbuf = vec![0f32; k * nkv * hd];
12804        let mut vbuf = vec![0f32; k * nkv * hd];
12805        if !crate::gpu_metal::kv_mirror_read_rows(
12806            self.mtp_kv_id(),
12807            Self::MTP_LAYER_BASE,
12808            nkv,
12809            hd,
12810            cpu_stored,
12811            k,
12812            &mut kbuf,
12813            &mut vbuf,
12814        ) {
12815            return Err(true);
12816        }
12817        for r in 0..k {
12818            m.kv.append(
12819                &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12820                &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12821                &[],
12822            );
12823        }
12824        Ok(ids)
12825    }
12826
12827    fn try_batch_graph_wgpu(
12828        &self,
12829        hiddens: &mut [f32],
12830        positions: &[usize],
12831        k: usize,
12832        spec: Option<crate::gpu::SpecTail<'_>>,
12833    ) -> crate::gpu::BatchGraphOutcome {
12834        self.try_batch_graph_wgpu_prefix(hiddens, positions, k, spec, None)
12835    }
12836
12837    /// `try_batch_graph_wgpu` with the device-prefix mode: `layers_run`
12838    /// Some lets a stack that does not fit run its leading layers (the
12839    /// token graph's prefix rule) and reports how many; `hiddens` then
12840    /// holds the boundary rows and the caller runs the rest on the host.
12841    fn try_batch_graph_wgpu_prefix(
12842        &self,
12843        hiddens: &mut [f32],
12844        positions: &[usize],
12845        k: usize,
12846        spec: Option<crate::gpu::SpecTail<'_>>,
12847        layers_run: Option<&mut usize>,
12848    ) -> crate::gpu::BatchGraphOutcome {
12849        let graph_end = match self.mimo_moe.graph_prefix_end() {
12850            Some(end) if end < self.num_layers => {
12851                if layers_run.is_none() || spec.is_some() || end == 0 {
12852                    return crate::gpu::BatchGraphOutcome::Declined;
12853                }
12854                end
12855            }
12856            _ => self.num_layers,
12857        };
12858        let _tb = std::time::Instant::now();
12859        let batch_debug = std::env::var_os("CMF_BATCH_DEBUG").is_some();
12860        if self.attn_softcap > 0.0 {
12861            return crate::gpu::BatchGraphOutcome::Declined; // capped scores: no graph kernel — CPU path
12862        }
12863        // Same attention contract as the token graph: per-layer geometry
12864        // rides `geom`, anything it cannot express declines by name.
12865        if let Some(reason) = self.wgpu_graph_attn_decline() {
12866            self.note_graph_decline("wgpu batch graph", reason);
12867            return crate::gpu::BatchGraphOutcome::Declined;
12868        }
12869        let nh = self.num_heads;
12870        let (nkv, hd, rd) = self.layer_geom(0);
12871        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
12872        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
12873            if let Some((m, i, kind, rs)) = t
12874                .graph_weight()
12875                .or_else(|| t.graph_weight_descriptor())
12876            {
12877                let name = &m.tensors[i].name;
12878                let prism = if crate::prism::is_inverse_embedding(m, name) {
12879                    crate::gpu::GraphPrismOp::InverseEmbedding
12880                } else if crate::prism::is_forward_weight(m, name) {
12881                    crate::gpu::GraphPrismOp::Forward
12882                } else {
12883                    crate::gpu::GraphPrismOp::None
12884                };
12885                return Some(crate::gpu::GraphW {
12886                    idx: i,
12887                    kind,
12888                    row_scale: rs,
12889                    data: &[],
12890                    prism,
12891                    affine: crate::prism::is_affine_target(m, name),
12892                });
12893            }
12894            if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
12895                eprintln!(
12896                    "batch graph: tensor has no graph descriptor/f32 fallback rows={} cols={}",
12897                    t.rows(),
12898                    t.cols()
12899                );
12900            }
12901            t.as_f32().map(|d| crate::gpu::GraphW {
12902                idx: 0,
12903                kind: 4,
12904                row_scale: &[],
12905                data: d,
12906                prism: crate::gpu::GraphPrismOp::None,
12907                affine: false,
12908            })
12909        }
12910        let built: Option<(
12911            Vec<crate::gpu::GraphLayer<'_>>,
12912            std::sync::Arc<cortiq_core::CmfModel>,
12913        )> = (|| {
12914            let mut layers = Vec::with_capacity(graph_end);
12915            let mut model = None;
12916            for li in 0..graph_end {
12917                let lw = &self.weights.layers[self.phys_layer(li)];
12918                // MoE routes per token, so its experts are encoded token by
12919                // token inside the batched submit while attention and the
12920                // projections stay GEMMs. Refusing MoE here is what left
12921                // prefill running one position at a time: 33 tok/s against
12922                // 54 on decode, i.e. reading the prompt was slower than
12923                // writing the answer.
12924                let gffn = match &lw.ffn {
12925                    FfnKind::Dense(d) if !d.segs.is_empty() => {
12926                        if batch_debug {
12927                            eprintln!("batch graph: dense segmented FFN at layer {li}");
12928                        }
12929                        return None;
12930                    }
12931                    FfnKind::Dense(d) => {
12932                        let Some(act) = d.act.graph_act() else {
12933                            if batch_debug {
12934                                eprintln!(
12935                                    "batch graph: dense FFN activation {:?} without a graph kernel at layer {li}",
12936                                    d.act
12937                                );
12938                            }
12939                            return None;
12940                        };
12941                        crate::gpu::GraphFfn::Dense {
12942                            gate: gw(&d.gate_proj)?,
12943                            up: gw(&d.up_proj)?,
12944                            down: gw(&d.down_proj)?,
12945                            act,
12946                        }
12947                    }
12948                    FfnKind::Moe(m) => {
12949                        // Adaptive τ and expert masks stay on the CPU path.
12950                        // Sigmoid scores, the selection bias, a routed scale
12951                        // ≠ 1 and an ungated shared expert (hy_v3) ride the
12952                        // same flags word as the token graph — before, this
12953                        // refusal sent every Hy-MT2-30B prompt to the chunked
12954                        // fallback (8 tok/s of ingest against 53 of decode).
12955                        if m.route_tau.is_some() || m.mask.is_some() {
12956                            return None;
12957                        }
12958                        // A shared expert rides as slot top_k (gated or
12959                        // not is a flag on the select kernel); without one
12960                        // (MiMo-V2, LFM2-MoE) the kernels run top_k slots.
12961                        let shared = m.shared.as_ref();
12962                        let has_shared = shared.is_some();
12963                        let shared_gated = matches!(shared, Some((_, Some(_))));
12964                        let sgate = match shared {
12965                            Some((_, Some(sg))) => gw(sg)?,
12966                            // Ungated or absent: the router plane stands in
12967                            // so the plumbing stays total; the kernel pins
12968                            // weight 1 or never reads it.
12969                            _ => gw(&m.router)?,
12970                        };
12971                        let router = gw(&m.router)?;
12972                        // The batch MoE kernels still consume raw per-token
12973                        // rows and do not carry the descriptor-aware Prism
12974                        // transform/affine bit for router or shared-gate
12975                        // planes.  Refuse rather than route an untransformed
12976                        // source activation.
12977                        if router.prism != crate::gpu::GraphPrismOp::None
12978                            || router.affine
12979                            || sgate.prism != crate::gpu::GraphPrismOp::None
12980                            || sgate.affine
12981                        {
12982                            return None;
12983                        }
12984                        let inter = m.experts.first()?.gate_proj.rows();
12985                        let mut experts = Vec::with_capacity(m.experts.len() + 1);
12986                        let mut q4tp: Option<bool> = None;
12987                        let mut gu_q2: Option<bool> = None;
12988                        for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
12989                            if !matches!(e.act, Act::Silu)
12990                                || e.gate_proj.rows() != inter
12991                                || e.up_proj.rows() != inter
12992                            {
12993                                return None;
12994                            }
12995                            // Same ladder as the token graph: q4t → q2tp
12996                            // (mixed profile: 2-bit gate/up over a q4tp
12997                            // down) → q4tp. Uniform across the layer.
12998                            let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
12999                                Some((mm, gi)) => (
13000                                    mm,
13001                                    gi,
13002                                    e.up_proj.mapped_q4t()?.1,
13003                                    e.down_proj.mapped_q4t()?.1,
13004                                    false,
13005                                    false,
13006                                ),
13007                                None => match e.gate_proj.mapped_q2tp() {
13008                                    Some((mm, gi)) => (
13009                                        mm,
13010                                        gi,
13011                                        e.up_proj.mapped_q2tp()?.1,
13012                                        e.down_proj.mapped_q4tp()?.1,
13013                                        true,
13014                                        true,
13015                                    ),
13016                                    None => {
13017                                        let (mm, gi) = e.gate_proj.mapped_q4tp()?;
13018                                        (
13019                                            mm,
13020                                            gi,
13021                                            e.up_proj.mapped_q4tp()?.1,
13022                                            e.down_proj.mapped_q4tp()?.1,
13023                                            true,
13024                                            false,
13025                                        )
13026                                    }
13027                                },
13028                            };
13029                            if *q4tp.get_or_insert(is_p) != is_p
13030                                || *gu_q2.get_or_insert(is_q2) != is_q2
13031                            {
13032                                return None;
13033                            }
13034                            if [gi, ui, di].into_iter().any(|idx| {
13035                                mm.tensors
13036                                    .get(idx)
13037                                    .is_some_and(|t| {
13038                                        crate::prism::is_forward_weight(mm, &t.name)
13039                                            || crate::prism::is_affine_target(mm, &t.name)
13040                                    })
13041                            }) {
13042                                return None;
13043                            }
13044                            model.get_or_insert_with(|| mm.clone());
13045                            experts.push((gi, ui, di));
13046                        }
13047                        crate::gpu::GraphFfn::Moe {
13048                            router,
13049                            shared_gate: sgate,
13050                            experts,
13051                            n_exp: m.experts.len(),
13052                            top_k: m.top_k,
13053                            inter,
13054                            norm_topk: m.norm_topk_prob,
13055                            q4tp: q4tp?,
13056                            gu_q2: gu_q2.unwrap_or(false),
13057                            sigmoid: m.router_sigmoid,
13058                            bias: m.expert_bias.as_deref(),
13059                            has_shared,
13060                            shared_gated,
13061                            route_scale: m.routed_scaling,
13062                        }
13063                    }
13064                    _ => return None,
13065                };
13066                let attn = match &lw.attn {
13067                    AttnKind::Full {
13068                        wq,
13069                        wk,
13070                        wv,
13071                        wo,
13072                        q_norm,
13073                        k_norm,
13074                        output_gate,
13075                        softplus_gate,
13076                        bias,
13077                    } => {
13078                        if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
13079                            if batch_debug {
13080                                eprintln!(
13081                                    "batch graph: unsupported Full attention gate at layer {li} softplus={} heads={}",
13082                                    softplus_gate.is_some(),
13083                                    self.attention_heads_per_layer.is_some()
13084                                );
13085                            }
13086                            return None;
13087                        }
13088                        let (m, _, _, _) = wq
13089                            .graph_weight()
13090                            .or_else(|| wq.graph_weight_descriptor())?;
13091                        model = Some(m.clone());
13092                        crate::gpu::GraphAttn::Full {
13093                            wq: gw(wq)?,
13094                            wk: gw(wk)?,
13095                            wv: gw(wv)?,
13096                            wo: gw(wo)?,
13097                            q_norm: q_norm.as_deref(),
13098                            k_norm: k_norm.as_deref(),
13099                            late_qk_norm: self.qk_norm_after_rope,
13100                            bias: bias
13101                                .as_ref()
13102                                .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
13103                            output_gate: *output_gate,
13104                            cpu_k: self.kv_cache.layers[li].k_heads(),
13105                            cpu_v: self.kv_cache.layers[li].v_heads(),
13106                            geom: self.graph_attn_geom(li),
13107                            // The batched graph has no head-gate arm; a
13108                            // gated layer is refused above.
13109                            head_gate: None,
13110                        }
13111                    }
13112                    AttnKind::LinearGdn(w) => {
13113                        let Some(cfg) = self.gdn_cfg else {
13114                            if batch_debug {
13115                                eprintln!("batch graph: no GDN config at layer {li}");
13116                            }
13117                            return None;
13118                        };
13119                        let (m, _, _, _) = w
13120                            .in_proj_qkv
13121                            .graph_weight()
13122                            .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
13123                        model = Some(m.clone());
13124                        crate::gpu::GraphAttn::Gdn {
13125                            qkv: gw(&w.in_proj_qkv)?,
13126                            z: gw(&w.in_proj_z)?,
13127                            a: gw(&w.in_proj_a)?,
13128                            b: gw(&w.in_proj_b)?,
13129                            out: gw(&w.out_proj)?,
13130                            conv1d: &w.conv1d,
13131                            a_log: &w.a_log,
13132                            dt_bias: &w.dt_bias,
13133                            norm: &w.norm,
13134                            nv: cfg.num_v_heads,
13135                            nk: cfg.num_k_heads,
13136                            dk: cfg.key_head_dim,
13137                            dv: cfg.value_head_dim,
13138                            kk: cfg.conv_kernel,
13139                            cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
13140                        }
13141                    }
13142                    _ => return None,
13143                };
13144                layers.push(crate::gpu::GraphLayer {
13145                    input_norm: &lw.input_norm,
13146                    attn,
13147                    post_norm: &lw.post_norm,
13148                    ffn: gffn,
13149                });
13150            }
13151            Some((layers, model?))
13152        })();
13153        let Some((layers, model)) = built else {
13154            {
13155                use std::sync::atomic::{AtomicBool, Ordering};
13156                static SAID: AtomicBool = AtomicBool::new(false);
13157                if !SAID.swap(true, Ordering::Relaxed) {
13158                    tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
13159                }
13160            }
13161            return crate::gpu::BatchGraphOutcome::Declined;
13162        };
13163        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
13164            eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
13165        }
13166        crate::gpu::forward_batch_graph(
13167            &model,
13168            self.graph_kv_id,
13169            &layers,
13170            &self.inv_freq,
13171            hiddens,
13172            nh,
13173            nkv,
13174            hd,
13175            rd,
13176            self.hidden_size,
13177            self.intermediate_size,
13178            positions,
13179            self.kv_cache.max_seq_len,
13180            gemma,
13181            self.rms_eps as f32,
13182            self.attn_scale,
13183            k,
13184            &(0..graph_end)
13185                .map(|li| self.kv_cache.layers[self.phys_layer(li)].o1_views())
13186                .collect::<Vec<_>>(),
13187            self.o1_epoch,
13188            spec,
13189            layers_run,
13190        )
13191    }
13192
13193    /// Same, stopping after layer `upto` inclusive (routing probe φ).
13194    /// `CMF_DSV4_DRAFT_PROBE=1` — grade the draft against what the trunk goes on
13195    /// to produce. Off by default; it runs a whole draft per decoded token.
13196    fn draft_probe() -> bool {
13197        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13198        *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
13199    }
13200
13201    /// `CMF_DSV4_DRAFT_PROBE=1`: measure how much of the draft the trunk
13202    /// would have agreed with, WITHOUT verifying or rolling anything back.
13203    ///
13204    /// The number this produces decides the whole speculation design — at
13205    /// acceptance a, a block of B positions yields 1 + a + a² + ... tokens
13206    /// per trunk pass — so it is worth measuring before any of the machinery
13207    /// that would exploit it exists. Each draft is parked with the position
13208    /// it was made at, and graded as the real tokens arrive.
13209    /// `CMF_DSV4_SPEC=1` — the DeepSeek-V4 speculative decode: draft five
13210    /// on the card, verify them in one batched trunk pass, commit the
13211    /// accepted prefix, roll the rest back.
13212    #[cfg(feature = "gpu")]
13213    fn dsv4_spec_on() -> bool {
13214        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13215        *ON.get_or_init(|| {
13216            // Test-only runtime gate: model loading still performs the same
13217            // reservation and trunk packing, which gives rollback parity a
13218            // topology-identical non-speculative control arm.
13219            if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
13220                return v != "0";
13221            }
13222            // An explicit value is a diagnostic force/escape hatch.  With no
13223            // knob, speculation is eligible only when model loading reserved
13224            // its bounded pack.  On small q4tp cards the geometric reserve
13225            // gate deliberately leaves this at zero: trying to build DSpark
13226            // after the exact trunk filled VRAM is both slower and a device
13227            // OOM (measured on A40).
13228            std::env::var("CMF_DSV4_SPEC")
13229                .map(|v| v != "0")
13230                .unwrap_or_else(|_| {
13231                    crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
13232                })
13233        })
13234    }
13235
13236    /// One speculative round at the decode tip. `t_next` is the token the
13237    /// sampler just committed for `next_pos`. Returns the EXTRA accepted
13238    /// tokens (possibly none) and the new position, with `graph_logits`
13239    /// left holding the last accepted position's logits — exactly what the
13240    /// loop top expects. `None` means "speculate not this round": nothing
13241    /// was committed, the caller forwards normally.
13242    #[cfg(feature = "gpu")]
13243    fn dsv4_spec_step(
13244        &mut self,
13245        tip_token: u32,
13246        t_next: u32,
13247        next_pos: usize,
13248        max_extra: usize,
13249        drafted: &mut usize,
13250        accepted_ctr: &mut usize,
13251    ) -> Option<(Vec<u32>, usize)> {
13252        let t_all = std::time::Instant::now();
13253        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13254            thread_local! {
13255                static LAST: std::cell::Cell<Option<std::time::Instant>> =
13256                    const { std::cell::Cell::new(None) };
13257            }
13258            LAST.with(|l| {
13259                if let Some(prev) = l.get() {
13260                    eprintln!(
13261                        "между раундами {:.1} мс",
13262                        prev.elapsed().as_secs_f64() * 1e3
13263                    );
13264                }
13265                l.set(Some(std::time::Instant::now()));
13266            });
13267        }
13268        if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13269            eprintln!("spec_step: вход pos={next_pos}");
13270        }
13271        let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
13272        let cfg = self.dsv4.as_ref().map(|b| b.2)?;
13273        // The draft state and its capture, armed exactly as the probe does.
13274        if self.dspark.is_none() {
13275            let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13276            if t.is_empty() {
13277                return None;
13278            }
13279            crate::dsv4::dspark_arm(&t, cfg.dim);
13280            self.dspark = Some(crate::dsv4::DsparkState::new(
13281                self.dsv4_mtp.len(),
13282                &cfg,
13283                t.len(),
13284            ));
13285        }
13286        let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13287        let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
13288        if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13289            eprintln!("spec_step: пак не построился (targets {targets:?})");
13290        }
13291        let pack = pack?;
13292        let block = crate::dsv4::dspark_block();
13293        let b_box = self.dsv4.as_mut()?;
13294        let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
13295        let ds = self.dspark.as_mut()?;
13296        // The tip's captures: either this token ran on a normal path that
13297        // filled the thread-local, or the previous spec round left them.
13298        let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
13299        if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
13300            if dbg {
13301                eprintln!("spec_step: нет захвата");
13302            }
13303            return None;
13304        }
13305        ds.have_hidden = true;
13306        let tip_pos = next_pos.checked_sub(1)?;
13307        let draft_started = std::time::Instant::now();
13308        let mut conf = Vec::new();
13309        let props = crate::dsv4::dspark_draft_gpu(
13310            g,
13311            &self.dsv4_mtp,
13312            &cfg,
13313            ds,
13314            pack,
13315            st.kv_id,
13316            tip_token,
13317            tip_pos,
13318            self.pool.as_deref(),
13319            &mut conf,
13320        );
13321        self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13322        *drafted += block;
13323        if props.is_empty() || props[0] != t_next {
13324            if dbg {
13325                eprintln!(
13326                    "spec_step: черновик {} (props0={:?} t_next={t_next})",
13327                    if props.is_empty() {
13328                        "пуст"
13329                    } else {
13330                        "мимо"
13331                    },
13332                    props.first()
13333                );
13334            }
13335            return None;
13336        }
13337        // `fed[0]` is `t_next`, which the outer loop has already committed;
13338        // only `fed[1..]` become additional output tokens. Cap the verify
13339        // transaction itself to the caller's remaining output budget instead
13340        // of merely truncating the returned vector: otherwise the KV/state
13341        // would advance past `max_tokens` and a 64-token request could return
13342        // 66 tokens (and poison a reused session with two invisible steps).
13343        let mut k_verify = crate::dsv4::dspark_verify_k()
13344            .min(props.len())
13345            .min(max_extra.saturating_add(1));
13346        // Adaptive depth: positions the draft itself doubts are paid for on
13347        // every verify and delivered almost never (natural-text survival
13348        // [.67 .50 .29 .08 .04]). `CMF_DSPARK_CONF_MIN=p` trims the fed
13349        // prefix at the first proposal whose confidence drops below p; on
13350        // predictable text the confidences stay high and nothing changes.
13351        let conf_min = {
13352            static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
13353            *M.get_or_init(|| {
13354                std::env::var("CMF_DSPARK_CONF_MIN")
13355                    .ok()
13356                    .and_then(|v| v.parse().ok())
13357                    .unwrap_or(0.0)
13358            })
13359        };
13360        if conf_min > 0.0 && conf.len() >= props.len() {
13361            let mut keep = 1usize;
13362            while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
13363                keep += 1;
13364            }
13365            k_verify = k_verify.min(keep.max(2));
13366        }
13367        if k_verify < 2 {
13368            return None;
13369        }
13370        let mut fed = Vec::with_capacity(k_verify);
13371        fed.push(t_next);
13372        fed.extend_from_slice(&props[1..k_verify]);
13373        let mut argmax = Vec::new();
13374        let mut logits_all = Vec::new();
13375        let mut walked = Vec::new();
13376        let txn = crate::dsv4::dsv4_verify_chunk(
13377            g,
13378            layers,
13379            &cfg,
13380            st,
13381            &fed,
13382            next_pos,
13383            &self.inv_freq,
13384            self.pool.as_deref(),
13385            &targets,
13386            &mut argmax,
13387            &mut logits_all,
13388            &mut walked,
13389        );
13390        if txn.is_none() && dbg {
13391            eprintln!("spec_step: verify отказал");
13392        }
13393        let txn = txn?;
13394        let spec_gpu_end = txn.gpu_end;
13395        let b = fed.len();
13396        let mut accepted = 1usize;
13397        while accepted < b && fed[accepted] == argmax[accepted - 1] {
13398            accepted += 1;
13399        }
13400        // `CMF_DSV4_SPEC_FORCE_REJECT=1` — accept nothing beyond the known
13401        // token, every round: the pure rollback exerciser. The output must
13402        // stay byte-identical to the plain walk; anything else is a
13403        // transaction bug, isolated from the acceptance logic.
13404        if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
13405            accepted = 1;
13406        }
13407        if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
13408            eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
13409        }
13410        let t_fin = std::time::Instant::now();
13411        if !crate::dsv4::dsv4_spec_finish(
13412            g,
13413            layers,
13414            &cfg,
13415            st,
13416            txn,
13417            accepted,
13418            &fed,
13419            &self.inv_freq,
13420            self.pool.as_deref(),
13421        ) {
13422            tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
13423            return None;
13424        }
13425        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13426            eprintln!(
13427                "finish(k={accepted}): {:.1} мс",
13428                t_fin.elapsed().as_secs_f64() * 1e3
13429            );
13430        }
13431        *accepted_ctr += accepted - 1;
13432        // Captures per accepted token: device targets photographed by the
13433        // batch, host targets from the verify's own walk. The last one
13434        // becomes the new tip's draft input; every one owes the ring an
13435        // entry for its position.
13436        let (hc, dim) = (cfg.hc_mult, cfg.dim);
13437        // Complete-chain layers are photographed by the fused submission;
13438        // partial device layers overwrite that slot after exact host cold-
13439        // expert correction.  Thus every target in the contiguous device
13440        // prefix has a valid per-token capture.
13441        let dev_caps: Vec<usize> = targets
13442            .iter()
13443            .copied()
13444            .filter(|&t| t < spec_gpu_end)
13445            .collect();
13446        let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
13447        if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
13448            return None;
13449        }
13450        for t in 0..accepted {
13451            let tip = t + 1 == accepted;
13452            for (slot, &tl) in targets.iter().enumerate() {
13453                if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
13454                    let lo = (di * b + t) * hc * dim;
13455                    crate::dsv4::dspark_capture(
13456                        &caps_all[lo..lo + hc * dim],
13457                        &cfg,
13458                        slot,
13459                        &mut ds.main_hidden,
13460                    );
13461                } else if tip
13462                    && crate::dsv4::dspark_peek_slot(slot, dim, {
13463                        let lo = slot * dim;
13464                        &mut ds.main_hidden[lo..lo + dim]
13465                    })
13466                {
13467                    // The tip's host-layer captures are the walk's own
13468                    // per-layer notes — exact. (The walk that ran last ended
13469                    // on exactly this token, on both the accept-all and the
13470                    // rollback path.)
13471                } else {
13472                    // Intermediate tokens: the post-tail state stands in for
13473                    // the per-layer capture on host targets below the last
13474                    // layer. Ring-entry quality only; the tip is exact.
13475                    crate::dsv4::dspark_capture(
13476                        &walked[t * hc * dim..(t + 1) * hc * dim],
13477                        &cfg,
13478                        slot,
13479                        &mut ds.main_hidden,
13480                    );
13481                }
13482            }
13483            crate::dsv4::dspark_ring_append(
13484                g,
13485                &self.dsv4_mtp,
13486                &cfg,
13487                ds,
13488                next_pos + t,
13489                self.pool.as_deref(),
13490            );
13491        }
13492        let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
13493        self.graph_logits = Some(row);
13494        // The speculative loop never runs the probe, so the trunk tally has
13495        // no other place to cycle. Armed only when someone asked for the
13496        // dump; the host tail is the only tallying path here, which is
13497        // precisely the population a partial pack would serve.
13498        if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
13499            crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
13500            crate::dsv4::pick_tally_arm();
13501        }
13502        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13503            eprintln!(
13504                "spec_step total {:.1} мс (k={accepted})",
13505                t_all.elapsed().as_secs_f64() * 1e3
13506            );
13507        }
13508        Some((fed[1..accepted].to_vec(), next_pos + accepted))
13509    }
13510
13511    fn dspark_probe(&mut self, position: usize, token_id: u32) {
13512        if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
13513            return;
13514        }
13515        // What the trunk just routed to, for this token.
13516        let trunk_now = crate::dsv4::pick_tally_take();
13517        crate::dsv4::trunk_freq_note(&trunk_now);
13518        if !trunk_now.is_empty() {
13519            self.dspark_trunk_picks.push(trunk_now);
13520            let keep = crate::dsv4::dspark_block();
13521            if self.dspark_trunk_picks.len() > keep {
13522                self.dspark_trunk_picks.remove(0);
13523            }
13524        }
13525        // Grade whatever is waiting: the token just decoded sits at
13526        // `position`, so it answers the draft made at `position - 1 - i`.
13527        for p in std::mem::take(&mut self.dspark_pending) {
13528            let Some(i) = position.checked_sub(p.0 + 1) else {
13529                continue;
13530            };
13531            let mut p = p;
13532            if i < p.1.len() {
13533                if p.2 && p.1[i] == token_id {
13534                    p.3 = i + 1;
13535                } else {
13536                    p.2 = false;
13537                }
13538                if i + 1 < p.1.len() {
13539                    self.dspark_pending.push(p);
13540                    continue;
13541                }
13542            }
13543            self.dspark_hist.push(p.3);
13544            self.dspark_real.push(token_id);
13545        }
13546        let Some(b) = &mut self.dsv4 else { return };
13547        let (g, layers, cfg) = (&b.0, &b.1, b.2);
13548        let n_layers = layers.len();
13549        if self.dspark.is_none() {
13550            let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13551            if t.is_empty() {
13552                return;
13553            }
13554            eprintln!(
13555                "DSpark: захват со слоёв {t:?}, блок {}",
13556                crate::dsv4::dspark_block()
13557            );
13558            crate::dsv4::dspark_arm(&t, cfg.dim);
13559            self.dspark = Some(crate::dsv4::DsparkState::new(
13560                self.dsv4_mtp.len(),
13561                &cfg,
13562                t.len(),
13563            ));
13564        }
13565        let ds = self.dspark.as_mut().unwrap();
13566        if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
13567            return; // this token ran on a path that captures nothing
13568        }
13569        let mut conf = Vec::new();
13570        crate::dsv4::pick_tally_arm();
13571        // The trunk has already consumed the adaptive VRAM budget. Until the
13572        // draft owns an explicit bounded device pack, its tensors are an
13573        // out-of-core CPU/disk tier by contract: never let per-op probes try
13574        // to squeeze another multi-gigabyte MTP expert cache onto the card.
13575        let draft_started = std::time::Instant::now();
13576        #[cfg(feature = "gpu")]
13577        let gpu_draft = crate::dsv4::dspark_gpu_on();
13578        #[cfg(not(feature = "gpu"))]
13579        let gpu_draft = false;
13580        let props = if gpu_draft {
13581            #[cfg(feature = "gpu")]
13582            {
13583                let kv_id = b.3.kv_id;
13584                match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
13585                    Some(pk) => crate::dsv4::dspark_draft_gpu(
13586                        g,
13587                        &self.dsv4_mtp,
13588                        &cfg,
13589                        ds,
13590                        pk,
13591                        kv_id,
13592                        token_id,
13593                        position,
13594                        self.pool.as_deref(),
13595                        &mut conf,
13596                    ),
13597                    None => Vec::new(),
13598                }
13599            }
13600            #[cfg(not(feature = "gpu"))]
13601            Vec::new()
13602        } else {
13603            crate::gpu::cpu_scope(|| {
13604                crate::dsv4::dspark_draft(
13605                    g,
13606                    &self.dsv4_mtp,
13607                    &cfg,
13608                    ds,
13609                    token_id,
13610                    position,
13611                    self.pool.as_deref(),
13612                    &mut conf,
13613                )
13614            })
13615        };
13616        self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13617        let draft_picks = crate::dsv4::pick_tally_take();
13618        crate::dsv4::dspark_freq_note(&draft_picks);
13619        // Re-arm for the NEXT trunk token; the probe runs after the forward,
13620        // so this is the only place that can.
13621        crate::dsv4::pick_tally_arm();
13622        if !props.is_empty() {
13623            // Two ratios, side by side: what a batched verify over the trunk
13624            // would read against what it asks for, and the same for the
13625            // draft's three stages. Near 1.0 means a batch amortises nothing.
13626            let (tu, tt) = {
13627                let flat: Vec<(usize, Vec<usize>)> = self
13628                    .dspark_trunk_picks
13629                    .iter()
13630                    .flat_map(|v| v.iter().cloned())
13631                    .collect();
13632                // Per layer, across the window of tokens.
13633                let mut per: std::collections::HashMap<usize, Vec<usize>> =
13634                    std::collections::HashMap::new();
13635                for (li, picks) in flat {
13636                    per.entry(li).or_default().extend(picks);
13637                }
13638                let n = per.len().max(1);
13639                let mut u = 0usize;
13640                let mut t = 0usize;
13641                for (_, v) in per {
13642                    t += v.len();
13643                    u += v.iter().collect::<std::collections::HashSet<_>>().len();
13644                }
13645                (u / n, t / n)
13646            };
13647            let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
13648            self.dspark_exp.push((tu, tt, du, dt));
13649            self.dspark_pending.push((position, props, true, 0));
13650        }
13651        if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
13652            let n = self.dspark_hist.len() as f32;
13653            let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
13654            let block = crate::dsv4::dspark_block();
13655            let mut at = vec![0usize; block + 1];
13656            for &k in &self.dspark_hist {
13657                at[k] += 1;
13658            }
13659            // Prefix survival: S_i = P(the first i positions all held).
13660            let mut surv = Vec::with_capacity(block);
13661            for i in 1..=block {
13662                let k = at[i..].iter().sum::<usize>() as f32 / n;
13663                surv.push(format!("{k:.2}"));
13664            }
13665            let distinct = self
13666                .dspark_real
13667                .iter()
13668                .collect::<std::collections::HashSet<_>>()
13669                .len();
13670            let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
13671                (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
13672            });
13673            let m = self.dspark_exp.len().max(1);
13674            eprintln!(
13675                "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
13676                 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
13677                self.dspark_hist.len(),
13678                mean + 1.0,
13679                surv.join(" ")
13680            );
13681            eprintln!(
13682                "DSpark: разных токенов {distinct} из {} (вырожденность), \
13683                 эксперты ствол {}/{} на слой за {block} токенов, \
13684                 черновик {}/{} за блок, draft {:.2} мс/блок",
13685                self.dspark_real.len(),
13686                tu / m,
13687                tt / m,
13688                du / m,
13689                dt / m,
13690                self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
13691            );
13692        }
13693    }
13694
13695    fn forward_layers_upto(
13696        &mut self,
13697        hidden: &[f32],
13698        position: usize,
13699        task_mask: Option<&TaskMask>,
13700        upto: Option<usize>,
13701    ) -> Vec<f32> {
13702        // In-process multi-GPU: each segment runs pinned to its card,
13703        // and the only thing crossing the boundary is one hidden vector
13704        // that never leaves this address space. Same layer split the
13705        // network mode does, minus the second process, the socket, the
13706        // serialization and the dir_hash handshake.
13707        if let Some(plan) = self.gpu_plan.clone() {
13708            if upto.is_none() && plan.len() > 1 {
13709                let mut h = hidden.to_vec();
13710                for &(dev, from, upto_incl) in plan.iter() {
13711                    h = crate::gpu::with_device(dev, || {
13712                        self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
13713                    });
13714                }
13715                return h;
13716            }
13717        }
13718        self.forward_layers_span(hidden, position, task_mask, 0, upto)
13719    }
13720
13721    /// Split this pipeline's layer stack across local GPUs: segment i
13722    /// runs on `devices[i]`. Contiguous and even by layer count — the
13723    /// VRAM-weighted planner is the next step, and an uneven card pair
13724    /// is why it will be needed. `None` clears the plan.
13725    pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
13726        self.set_gpu_plan_at(devices, None)
13727    }
13728
13729    /// The same, with an explicit first boundary (`--peer-split`): card
13730    /// 0 takes layers `[0..at)`, the rest split what remains. Uneven
13731    /// cards, or an attention-heavy head, are why this knob exists.
13732    pub fn set_gpu_plan_at(
13733        &mut self,
13734        devices: Option<&[usize]>,
13735        at: Option<usize>,
13736    ) -> Result<(), String> {
13737        let Some(devs) = devices.filter(|d| d.len() > 1) else {
13738            self.gpu_plan = None;
13739            return Ok(());
13740        };
13741        self.split_supported()?;
13742        let n = self.num_layers;
13743        if devs.len() > n {
13744            return Err(format!("{} devices for {n} layers", devs.len()));
13745        }
13746        if let Some(k) = at {
13747            if k == 0 || k >= n {
13748                return Err(format!("split at {k}: the model has {n} layers"));
13749            }
13750            if devs.len() == 2 {
13751                self.gpu_plan = Some(std::sync::Arc::new(vec![
13752                    (devs[0], 0, k - 1),
13753                    (devs[1], k, n - 1),
13754                ]));
13755                return Ok(());
13756            }
13757            return Err(format!(
13758                "an explicit split point takes exactly 2 devices, got {}",
13759                devs.len()
13760            ));
13761        }
13762        let per = n.div_ceil(devs.len());
13763        let mut plan = Vec::with_capacity(devs.len());
13764        let mut from = 0usize;
13765        for &d in devs {
13766            if from >= n {
13767                break;
13768            }
13769            let upto = (from + per - 1).min(n - 1);
13770            plan.push((d, from, upto));
13771            from = upto + 1;
13772        }
13773        self.gpu_plan = Some(std::sync::Arc::new(plan));
13774        Ok(())
13775    }
13776
13777    /// The active in-process split, if any: (device, first layer, last).
13778    pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
13779        self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
13780    }
13781
13782    /// Layer span [from ..= upto] (upto None = last layer): the building
13783    /// block the network pipeline-split rides on. `from > 0` skips the
13784    /// arch escape hatches (the pub `forward_span` refuses those archs
13785    /// first) and the whole-token graph — the plain per-layer loop is
13786    /// the canonical executor for a partial stack.
13787    fn embryo_qtensor_f32(t: &QTensor) -> Vec<f32> {
13788        if let Some(x) = t.as_f32() {
13789            return x.to_vec();
13790        }
13791        let mut out = vec![0.0; t.rows() * t.cols()];
13792        for r in 0..t.rows() {
13793            t.row_f32(r, &mut out[r * t.cols()..(r + 1) * t.cols()]);
13794        }
13795        out
13796    }
13797
13798    fn embryo_resident_eligible(&self) -> bool {
13799        // One mixer family per file: vmf_phase (kind 0/1) or
13800        // gated_delta_net (kind 4); the anchors are full (2) or bounded (3).
13801        if (self.vmf_cfg.is_none() && self.gdn_cfg.is_none())
13802            || self.num_layers != self.physical_layers
13803            || self.loop_final_norm
13804            || self.weights.layers.len() != self.num_layers
13805            || self.head_clusters.is_none()
13806            || self.final_softcap.is_some()
13807            || self.logit_multiplier.is_some()
13808            || self.attn_softcap != 0.0
13809            || self.mtp.is_some()
13810            || self.g3n.is_some()
13811            || self.dsv4.is_some()
13812            || self.dsv41.is_some()
13813            || self.qwen4_exp.is_some()
13814            // Dynamic routing swaps FFN weights mid-sequence under the
13815            // packed graph; a blend has no single overlay. A STATIC skill
13816            // (`from_model_with_skill`) is fine: the pack reads the live
13817            // `weights.layers[*].ffn`, i.e. the skill's tensors, and a
13818            // later `set_active_skill` change drops the pack
13819            // (`invalidate_for_weight_change`).
13820            || self.dyn_router.is_some()
13821            || self.dyn_phi_layer.is_some()
13822            || self.dyn_blend_loaded
13823            || self.o1_cfg.is_some()
13824            || self.swa.is_some()
13825            || self.sliding_layers.is_some()
13826            || self.global_attn.is_some()
13827            || self.attention_heads_per_layer.is_some()
13828            || self.attn_v_norm
13829            || self
13830                .kv_cache
13831                .layers
13832                .iter()
13833                .any(|l| l.mode != crate::kv_cache::KvMode::F32)
13834            || self.rope_scale != 1.0
13835            || self.rope_scale_local != 1.0
13836            || self.attn_scale != 1.0 / (self.head_dim as f32).sqrt()
13837            || self.hidden_size == 0
13838            || self.hidden_size > 1024
13839            || self.intermediate_size > 1024
13840            || self.num_heads == 0
13841            || self.num_kv_heads == 0
13842            || self.num_heads % self.num_kv_heads != 0
13843            || self.num_heads.saturating_mul(self.head_dim) > 1024
13844            || self.num_kv_heads.saturating_mul(self.head_dim) > 1024
13845            || self.vocab_size == 0
13846            || self.kv_cache.max_seq_len == 0
13847            || self.rotary_dim == 0
13848            || self.rotary_dim > self.head_dim
13849            || self.rotary_dim % 2 != 0
13850            || self.inv_freq.len() < self.rotary_dim / 2
13851        {
13852            return false;
13853        }
13854        // The resident shader is deliberately an f32 profile.  Dequantizing
13855        // a Q4/Q8 tensor into the packed buffer would silently change the
13856        // operator relative to the CPU quantized path, so quantized CMFs
13857        // retain the exact ordinary executor instead of claiming parity.
13858        // Measured consequence (RTX PRO 4000, S4 bounded export requantized
13859        // with `cortiq requant --quant q4tp-quantize`): `eligible=false`,
13860        // the generic wgpu whole-token graph refuses too, and the per-op
13861        // path decodes at ~73 tok/s against ~200 tok/s on the CPU q4tp
13862        // path — a q4tp Embryo-O1 file is a CPU artifact today; the
13863        // resident graph serves the f32 export.
13864        if self.weights.lm_head.as_f32().is_none()
13865            || self.weights.embed_tokens.as_f32().is_none()
13866            || self.weights.lm_head.rows() < self.vocab_size
13867            || self.weights.lm_head.cols() != self.hidden_size
13868            || self.weights.embed_tokens.rows() < self.vocab_size
13869            || self.weights.embed_tokens.cols() != self.hidden_size
13870            || self.weights.final_norm.len() != self.hidden_size
13871        {
13872            return false;
13873        }
13874        if let Some(cfg) = self.vmf_cfg {
13875            if cfg.num_heads.saturating_mul(cfg.nphase) > 1024
13876                || cfg.num_heads.saturating_mul(cfg.nphase.saturating_add(1)) > 1024
13877                || cfg.num_heads.saturating_mul(cfg.value_head_dim) > 1024
13878                || cfg.state_len() == 0
13879            {
13880                return false;
13881            }
13882        }
13883        if let Some(g) = self.gdn_cfg {
13884            // The resident GDN kernels (gpu_wgpu.rs `embryo_core_gdn_*`):
13885            // fused projection ≤ 2048 rows, nv·dv ≤ 1024, dk ≤ 128 lanes,
13886            // dv ≤ 256 lanes in vec4 rows, SiLU output gate (the Embryo
13887            // export), same rms eps as the stack.
13888            if g.num_v_heads == 0
13889                || g.num_k_heads == 0
13890                || g.num_v_heads % g.num_k_heads != 0
13891                || g.key_head_dim == 0
13892                || g.key_head_dim > 128
13893                || g.value_head_dim == 0
13894                || g.value_head_dim > 256
13895                || g.value_head_dim % 4 != 0
13896                || g.conv_kernel == 0
13897                || g.num_v_heads > 512
13898                || g.num_v_heads.saturating_mul(g.value_head_dim) > 1024
13899                || g.conv_dim() > 2048
13900                || g.conv_dim() % 4 != 0
13901                || g.hidden_size != self.hidden_size
13902                || g.output_gate_sigmoid
13903                || g.rms_eps != self.rms_eps
13904                || g.state_len() == 0
13905            {
13906                return false;
13907            }
13908        }
13909        let mut full_seen = false;
13910        for lw in &self.weights.layers {
13911            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
13912                return false;
13913            }
13914            match &lw.attn {
13915                AttnKind::LinearGdn(w) => {
13916                    let Some(g) = self.gdn_cfg else {
13917                        return false;
13918                    };
13919                    let (nv, dv, kk) = (g.num_v_heads, g.value_head_dim, g.conv_kernel);
13920                    if w.in_proj_qkv.rows() != g.conv_dim()
13921                        || w.in_proj_qkv.cols() != self.hidden_size
13922                        || w.in_proj_qkv.as_f32().is_none()
13923                        || w.in_proj_z.rows() != nv * dv
13924                        || w.in_proj_z.cols() != self.hidden_size
13925                        || w.in_proj_z.as_f32().is_none()
13926                        || w.in_proj_a.rows() != nv
13927                        || w.in_proj_a.cols() != self.hidden_size
13928                        || w.in_proj_a.as_f32().is_none()
13929                        || w.in_proj_b.rows() != nv
13930                        || w.in_proj_b.cols() != self.hidden_size
13931                        || w.in_proj_b.as_f32().is_none()
13932                        || w.conv1d.len() != g.conv_dim() * kk
13933                        || w.a_log.len() != nv
13934                        || w.dt_bias.len() != nv
13935                        || w.norm.len() != dv
13936                        || w.out_proj.rows() != self.hidden_size
13937                        || w.out_proj.cols() != nv * dv
13938                        || w.out_proj.as_f32().is_none()
13939                    {
13940                        return false;
13941                    }
13942                }
13943                AttnKind::Linear(w) => {
13944                    let Some(cfg) = self.vmf_cfg else {
13945                        return false;
13946                    };
13947                    if w.thq.rows() != cfg.num_heads * cfg.nphase
13948                        || w.thq.cols() != self.hidden_size
13949                        || w.thq.as_f32().is_none()
13950                        || w.thk.rows() != cfg.num_heads * cfg.nphase
13951                        || w.thk.cols() != self.hidden_size
13952                        || w.thk.as_f32().is_none()
13953                        || w.v_proj.rows() != cfg.num_heads * cfg.value_head_dim
13954                        || w.v_proj.cols() != self.hidden_size
13955                        || w.v_proj.as_f32().is_none()
13956                        || w.out_proj.rows() != self.hidden_size
13957                        || w.out_proj.cols() != cfg.num_heads * cfg.value_head_dim
13958                        || w.out_proj.as_f32().is_none()
13959                        || w.decay.len() != cfg.num_heads * 2 * cfg.nphase
13960                    {
13961                        return false;
13962                    }
13963                    if let Some((kg, kb)) = &w.k_gate {
13964                        if kg.rows() != cfg.num_heads
13965                            || kg.cols() != self.hidden_size
13966                            || kg.as_f32().is_none()
13967                            || kb.len() != cfg.num_heads
13968                        {
13969                            return false;
13970                        }
13971                    }
13972                    if let Some(conv) = &w.conv {
13973                        if conv.len() % self.hidden_size != 0 || conv.len() / self.hidden_size < 2 {
13974                            return false;
13975                        }
13976                    }
13977                }
13978                AttnKind::Full {
13979                    wq,
13980                    wk,
13981                    wv,
13982                    wo,
13983                    q_norm,
13984                    k_norm,
13985                    output_gate,
13986                    softplus_gate,
13987                    bias,
13988                } => {
13989                    if full_seen
13990                        || q_norm.is_some()
13991                        || k_norm.is_some()
13992                        || *output_gate
13993                        || softplus_gate.is_some()
13994                        || bias.is_some()
13995                        || wq.as_f32().is_none()
13996                        || wk.as_f32().is_none()
13997                        || wv.as_f32().is_none()
13998                        || wo.as_f32().is_none()
13999                        || wq.rows() != self.num_heads * self.head_dim
14000                        || wk.rows() != self.num_kv_heads * self.head_dim
14001                        || wv.rows() != self.num_kv_heads * self.head_dim
14002                        || wq.cols() != self.hidden_size
14003                        || wk.cols() != self.hidden_size
14004                        || wv.cols() != self.hidden_size
14005                        || wo.rows() != self.hidden_size
14006                        || wo.cols() != self.num_heads * self.head_dim
14007                    {
14008                        return false;
14009                    }
14010                    full_seen = true;
14011                }
14012                AttnKind::Bounded(w) => {
14013                    // The resident bounded attend scores S + W lanes in one
14014                    // 256-lane chunk; the format caps S + W at 160.
14015                    let Some(ac) = self.anchor_core.as_ref() else {
14016                        return false;
14017                    };
14018                    if self.bounded_rope.is_none()
14019                        || w.window != ac.window
14020                        || w.sink != ac.sink
14021                        || w.window == 0
14022                        || w.window + w.sink > 256
14023                        || w.sink_k.len() != self.num_kv_heads * w.sink * self.head_dim
14024                        || w.sink_v.len() != self.num_kv_heads * w.sink * self.head_dim
14025                        || w.wq.as_f32().is_none()
14026                        || w.wk.as_f32().is_none()
14027                        || w.wv.as_f32().is_none()
14028                        || w.wo.as_f32().is_none()
14029                        || w.wq.rows() != self.num_heads * self.head_dim
14030                        || w.wk.rows() != self.num_kv_heads * self.head_dim
14031                        || w.wv.rows() != self.num_kv_heads * self.head_dim
14032                        || w.wq.cols() != self.hidden_size
14033                        || w.wk.cols() != self.hidden_size
14034                        || w.wv.cols() != self.hidden_size
14035                        || w.wo.rows() != self.hidden_size
14036                        || w.wo.cols() != self.num_heads * self.head_dim
14037                    {
14038                        return false;
14039                    }
14040                }
14041                _ => return false,
14042            }
14043            match &lw.ffn {
14044                FfnKind::Dense(d) => {
14045                    if d.act != Act::Silu
14046                        || !d.segs.is_empty()
14047                        || d.gate_proj.as_f32().is_none()
14048                        || d.up_proj.as_f32().is_none()
14049                        || d.down_proj.as_f32().is_none()
14050                        || d.gate_proj.rows() != self.intermediate_size
14051                        || d.gate_proj.cols() != self.hidden_size
14052                        || d.up_proj.rows() != self.intermediate_size
14053                        || d.up_proj.cols() != self.hidden_size
14054                        || d.down_proj.rows() != self.hidden_size
14055                        || d.down_proj.cols() != self.intermediate_size
14056                    {
14057                        return false;
14058                    }
14059                }
14060                FfnKind::Moe(m) => {
14061                    if m.resonance.is_none()
14062                        || m.top_k != 1
14063                        || m.router_sigmoid
14064                        || !m.norm_topk_prob
14065                        || m.expert_bias.is_some()
14066                        || m.routed_scaling != 1.0
14067                        || m.route_tau.is_some()
14068                        || m.shared.is_none()
14069                        || m.mask.is_some()
14070                        || m.per_expert_scale.is_some()
14071                        || m.router_input_norm
14072                        || m.experts.is_empty()
14073                        || m.experts.len() > 8
14074                    {
14075                        return false;
14076                    }
14077                    let r = m.resonance.as_ref().unwrap();
14078                    if r.mu.len() != m.experts.len() * self.hidden_size
14079                        || r.bias.len() != m.experts.len()
14080                        || r.u.len() != m.experts.len() * r.k * self.hidden_size
14081                        || r.k > 128
14082                    {
14083                        return false;
14084                    }
14085                    let Some((shared, gate)) = &m.shared else {
14086                        return false;
14087                    };
14088                    if gate.is_some() || shared.act != Act::Silu || !shared.segs.is_empty() {
14089                        return false;
14090                    }
14091                    if shared.gate_proj.as_f32().is_none()
14092                        || shared.up_proj.as_f32().is_none()
14093                        || shared.down_proj.as_f32().is_none()
14094                        || shared.gate_proj.rows() != self.intermediate_size
14095                        || shared.gate_proj.cols() != self.hidden_size
14096                        || shared.up_proj.rows() != self.intermediate_size
14097                        || shared.up_proj.cols() != self.hidden_size
14098                        || shared.down_proj.rows() != self.hidden_size
14099                        || shared.down_proj.cols() != self.intermediate_size
14100                    {
14101                        return false;
14102                    }
14103                    for e in &m.experts {
14104                        if e.act != Act::Silu
14105                            || !e.segs.is_empty()
14106                            || e.gate_proj.as_f32().is_none()
14107                            || e.up_proj.as_f32().is_none()
14108                            || e.down_proj.as_f32().is_none()
14109                            || e.gate_proj.rows() != self.intermediate_size
14110                            || e.gate_proj.cols() != self.hidden_size
14111                            || e.up_proj.rows() != self.intermediate_size
14112                            || e.up_proj.cols() != self.hidden_size
14113                            || e.down_proj.rows() != self.hidden_size
14114                            || e.down_proj.cols() != self.intermediate_size
14115                        {
14116                            return false;
14117                        }
14118                    }
14119                }
14120                FfnKind::DenseMoe(_) => return false,
14121            }
14122            if lw.input_norm.len() != self.hidden_size || lw.post_norm.len() != self.hidden_size {
14123                return false;
14124            }
14125        }
14126        if full_seen && self.anchor_core.is_some() {
14127            return false;
14128        }
14129        full_seen || self.num_layers > 0
14130    }
14131
14132    /// The resident Embryo graph is the owner of this pipeline's forward:
14133    /// the same gate `forward_layers_span` applies before handing a token
14134    /// to `forward_embryo_graph` (both graph phases on, the explicit
14135    /// opt-in, a wgpu device, no earlier refusal, an eligible stack).
14136    fn embryo_resident_wanted(&self) -> bool {
14137        crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14138            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14139            && matches!(
14140                std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14141                Ok("1") | Ok("parallel")
14142            )
14143            && crate::gpu::enabled_here()
14144            && !self.graph_refused()
14145            && self.embryo_resident_eligible()
14146    }
14147
14148    /// Chunked prefill on the resident graph: `ids` from `start` in
14149    /// chunks of `EMBRYO_CHUNK_MAX`, one submit each, projections/FFN as
14150    /// chunk GEMMs and the recurrent layers walked in time on the device.
14151    /// Returns the last position's logits when the whole span ran there.
14152    /// `None` = refused before any device work (the per-position path
14153    /// takes the span).  `CMF_EMBRYO_CHUNK=0` keeps the per-position
14154    /// prefill (A/B and the parity reference).
14155    fn embryo_prefill_chunked(&mut self, ids: &[u32], start: usize) -> Option<Vec<f32>> {
14156        if ids.len() < 2
14157            || std::env::var("CMF_EMBRYO_CHUNK").as_deref() == Ok("0")
14158            || !self.embryo_resident_wanted()
14159        {
14160            return None;
14161        }
14162        let model = self.ensure_embryo_graph()?;
14163        let cmax = std::env::var("CMF_EMBRYO_CHUNK")
14164            .ok()
14165            .and_then(|v| v.parse::<usize>().ok())
14166            .filter(|&v| v >= 1)
14167            .unwrap_or(crate::gpu::EMBRYO_CHUNK_MAX)
14168            .min(crate::gpu::EMBRYO_CHUNK_MAX);
14169        let hs = self.hidden_size;
14170        let n = ids.len();
14171        let mut pos = start;
14172        let mut last = None;
14173        let mut rows = Vec::with_capacity(cmax * hs);
14174        while pos < n {
14175            let end = (pos + cmax).min(n);
14176            rows.clear();
14177            for &id in &ids[pos..end] {
14178                rows.extend_from_slice(&self.embed_single(id));
14179            }
14180            let mut lg = Vec::new();
14181            if !crate::gpu::forward_embryo_graph_chunk(
14182                &model,
14183                self.graph_kv_id,
14184                &rows,
14185                pos,
14186                end - pos,
14187                &mut lg,
14188            ) {
14189                if pos == start {
14190                    if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14191                        eprintln!("embryo-dbg: chunk prefill refused at position {pos}");
14192                    }
14193                    return None;
14194                }
14195                // The device sequence advanced through the earlier chunks;
14196                // a host continuation would mix two owners of the state.
14197                // Fail this sequence alone, leaving no stale state or key.
14198                self.kv_cache.clear();
14199                self.clear_history();
14200                crate::gpu::graph_kv_reset(self.graph_kv_id);
14201                panic!(
14202                    "resident Embryo chunk prefill refused at position {pos}; refusing an in-flight host fallback"
14203                );
14204            }
14205            last = Some(lg);
14206            pos = end;
14207        }
14208        if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14209            eprintln!(
14210                "embryo-dbg: chunk prefill {} ids from {start} in {} submits (chunk {cmax})",
14211                n - start,
14212                (n - start).div_ceil(cmax)
14213            );
14214        }
14215        last
14216    }
14217
14218    fn ensure_embryo_graph(&mut self) -> Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>> {
14219        if self.embryo_graph.is_none() && self.embryo_resident_eligible() {
14220            const UMAX: u32 = u32::MAX;
14221            const HEADER: usize = crate::gpu::EMBRYO_META_HEADER;
14222            const REC: usize = 64;
14223            struct Pack {
14224                data: Vec<f32>,
14225            }
14226            impl Pack {
14227                fn put(&mut self, x: &[f32]) -> u32 {
14228                    if x.is_empty() {
14229                        return u32::MAX;
14230                    }
14231                    let off = self.data.len();
14232                    self.data.extend_from_slice(x);
14233                    off as u32
14234                }
14235            }
14236            // The mixer family of the file: vmf_phase geometry fills the
14237            // phase header words, gated_delta_net fills words 24..29.  A
14238            // file has exactly one linear core, so at most one is live.
14239            let vmf = self.vmf_cfg;
14240            let gdn = self.gdn_cfg;
14241            let mut pack = Pack { data: Vec::new() };
14242            let mut meta = vec![0u32; HEADER];
14243            meta[0] = self.hidden_size as u32;
14244            meta[1] = self.intermediate_size as u32;
14245            meta[2] = self.vocab_size as u32;
14246            meta[3] = self.num_layers as u32;
14247            meta[4] = vmf.map(|c| c.num_heads).unwrap_or(0) as u32;
14248            meta[5] = vmf.map(|c| c.nphase).unwrap_or(0) as u32;
14249            meta[6] = vmf.map(|c| c.value_head_dim).unwrap_or(0) as u32;
14250            if let Some(g) = gdn {
14251                meta[24] = g.num_v_heads as u32;
14252                meta[25] = g.num_k_heads as u32;
14253                meta[26] = g.key_head_dim as u32;
14254                meta[27] = g.value_head_dim as u32;
14255                meta[28] = g.conv_kernel as u32;
14256                meta[29] = g.conv_dim() as u32;
14257            }
14258            meta[7] = self.num_heads as u32;
14259            meta[8] = self.num_kv_heads as u32;
14260            meta[9] = self.head_dim as u32;
14261            meta[10] = self.kv_cache.max_seq_len as u32;
14262            let clusters = self.head_clusters.as_ref().unwrap();
14263            let cluster_count = clusters.len() / self.hidden_size;
14264            if clusters.len() % self.hidden_size != 0
14265                || cluster_count == 0
14266                || cluster_count > 1024
14267                || self.vocab_size % cluster_count != 0
14268                || self.weights.lm_head.rows() < self.vocab_size
14269                || self.weights.final_norm.len() != self.hidden_size
14270            {
14271                return None;
14272            }
14273            meta[11] = cluster_count as u32;
14274            meta[12] = (self.vocab_size / cluster_count) as u32;
14275            meta[13] = vmf.map(|c| c.state_len()).unwrap_or(0) as u32;
14276            meta[16] = self.rotary_dim as u32;
14277            meta[17] = matches!(self.norm_style, NormStyle::Gemma) as u32;
14278            meta[19] = (self.rms_eps as f32).to_bits();
14279            let max_conv = self
14280                .weights
14281                .layers
14282                .iter()
14283                .filter_map(|lw| match &lw.attn {
14284                    AttnKind::Linear(w) => w.conv.as_ref().map(|v| v.len() / self.hidden_size),
14285                    _ => None,
14286                })
14287                .max()
14288                .unwrap_or(1);
14289            // One state slot per recurrent layer: the phase state plus its
14290            // hidden-wide conv ring, or the GDN record `[conv ring | S]`
14291            // (`GdnCfg::state_len`).  A file carries one mixer family, so
14292            // the stride is exactly that family's record and the device
14293            // state buffer equals the header's recurrent bytes.
14294            let phase_stride = vmf
14295                .map(|c| c.state_len() + max_conv.saturating_sub(1) * self.hidden_size)
14296                .unwrap_or(0);
14297            let gdn_stride = gdn.map(|g| g.state_len()).unwrap_or(0);
14298            let state_stride = phase_stride.max(gdn_stride);
14299            // Bounded genome: the KV plane of an anchor is its ring
14300            // `[kvh][W][hd]` K + V, and only anchors own one.  Legacy full
14301            // anchors keep the `max_seq` planes indexed by layer.
14302            let bounded = self.anchor_core.clone();
14303            let (anchor_window, anchor_sink) = bounded
14304                .as_ref()
14305                .map(|ac| (ac.window, ac.sink))
14306                .unwrap_or((0, 0));
14307            let kv_stride = if bounded.is_some() {
14308                2usize
14309                    .saturating_mul(self.num_kv_heads)
14310                    .saturating_mul(anchor_window)
14311                    .saturating_mul(self.head_dim)
14312            } else {
14313                2usize
14314                    .saturating_mul(self.num_kv_heads)
14315                    .saturating_mul(self.kv_cache.max_seq_len)
14316                    .saturating_mul(self.head_dim)
14317            };
14318            meta[14] = state_stride as u32;
14319            meta[15] = kv_stride as u32;
14320            meta[18] = anchor_window as u32;
14321            meta[20] = anchor_sink as u32;
14322            meta[21] = match &self.bounded_rope {
14323                Some(rope) => {
14324                    // [W][half] cos then [W][half] sin, one contiguous table.
14325                    let off = pack.put(&rope.cos);
14326                    let _ = pack.put(&rope.sin);
14327                    off
14328                }
14329                None => UMAX,
14330            };
14331            let mut full_seen = false;
14332            let mut bounded_seen = 0usize;
14333            // Recurrent state slots belong to mixer layers only (phase or
14334            // GDN): an anchor owns a ring, not a state stride, so the
14335            // device state buffer is exactly the header's recurrent bytes.
14336            let mut phase_seen = 0usize;
14337            let mut gdn_seen = 0usize;
14338            for (li, lw) in self.weights.layers.iter().enumerate() {
14339                let base = meta.len();
14340                meta.resize(base + REC, UMAX);
14341                meta[base] = match &lw.attn {
14342                    AttnKind::Linear(w) if w.phase_delta => 1,
14343                    AttnKind::Linear(_) => 0,
14344                    AttnKind::Full { .. } => 2,
14345                    AttnKind::Bounded(_) => 3,
14346                    AttnKind::LinearGdn(_) => 4,
14347                    _ => UMAX,
14348                };
14349                meta[base + 1] = pack.put(&lw.input_norm);
14350                meta[base + 2] = pack.put(&lw.post_norm);
14351                meta[base + 25] = match &lw.attn {
14352                    AttnKind::Linear(_) => {
14353                        let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14354                        phase_seen += 1;
14355                        off
14356                    }
14357                    AttnKind::LinearGdn(_) => {
14358                        let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14359                        gdn_seen += 1;
14360                        off
14361                    }
14362                    _ => UMAX,
14363                };
14364                match &lw.attn {
14365                    AttnKind::LinearGdn(w) => {
14366                        // Layer record words 56..63 + 29, as the resident
14367                        // kernels read them (gpu_wgpu.rs `embryo_core_gdn_*`).
14368                        meta[base + 56] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_qkv));
14369                        meta[base + 57] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_z));
14370                        meta[base + 58] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_a));
14371                        meta[base + 59] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_b));
14372                        meta[base + 60] = pack.put(&w.conv1d);
14373                        meta[base + 61] = pack.put(&w.a_log);
14374                        meta[base + 62] = pack.put(&w.dt_bias);
14375                        meta[base + 63] = pack.put(&w.norm);
14376                        meta[base + 29] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14377                        meta[base + 24] = 0;
14378                    }
14379                    AttnKind::Linear(w) => {
14380                        meta[base + 3] = pack.put(&Self::embryo_qtensor_f32(&w.thq));
14381                        meta[base + 4] = pack.put(&Self::embryo_qtensor_f32(&w.thk));
14382                        meta[base + 5] = pack.put(&Self::embryo_qtensor_f32(&w.v_proj));
14383                        meta[base + 6] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14384                        let decay: Vec<f32> = w.decay.iter().map(|&x| x as f32).collect();
14385                        meta[base + 7] = pack.put(&decay);
14386                        if let Some((kg, kb)) = &w.k_gate {
14387                            meta[base + 8] = pack.put(&Self::embryo_qtensor_f32(kg));
14388                            meta[base + 9] = pack.put(kb);
14389                        }
14390                        if let Some(conv) = &w.conv {
14391                            meta[base + 10] = pack.put(conv);
14392                            meta[base + 24] = (conv.len() / self.hidden_size) as u32;
14393                        } else {
14394                            meta[base + 24] = 0;
14395                        }
14396                    }
14397                    AttnKind::Full { wq, wk, wv, wo, .. } => {
14398                        full_seen = true;
14399                        meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(wq));
14400                        meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(wk));
14401                        meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(wv));
14402                        meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(wo));
14403                        meta[base + 26] = (li * kv_stride) as u32;
14404                    }
14405                    AttnKind::Bounded(w) => {
14406                        meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(&w.wq));
14407                        meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(&w.wk));
14408                        meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(&w.wv));
14409                        meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(&w.wo));
14410                        // Ring slot of this anchor (anchors only, packed).
14411                        meta[base + 26] = (bounded_seen * kv_stride) as u32;
14412                        meta[base + 27] = pack.put(&w.sink_k);
14413                        meta[base + 28] = pack.put(&w.sink_v);
14414                        bounded_seen += 1;
14415                    }
14416                    _ => return None,
14417                }
14418                match &lw.ffn {
14419                    FfnKind::Dense(d) => {
14420                        meta[base + 15] = 0;
14421                        meta[base + 16] = 0;
14422                        meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&d.gate_proj));
14423                        meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&d.up_proj));
14424                        meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&d.down_proj));
14425                    }
14426                    FfnKind::Moe(m) => {
14427                        let r = m.resonance.as_ref().unwrap();
14428                        let (shared, _) = m.shared.as_ref().unwrap();
14429                        meta[base + 15] = 1;
14430                        meta[base + 16] = m.experts.len() as u32;
14431                        meta[base + 17] = pack.put(&r.mu);
14432                        meta[base + 18] = pack.put(&r.u);
14433                        meta[base + 19] = pack.put(&r.bias);
14434                        meta[base + 20] = r.k as u32;
14435                        // Word 30: the growth shell as the runtime applies
14436                        // it now (`+inf` on trunk rows, the stored finite
14437                        // shell on grown rows, all `+inf` under
14438                        // `CMF_GROWTH_SHELL=off`), followed by one `−∞`
14439                        // sentinel at index E the kernel writes as the
14440                        // score of an expert outside its shell — WGSL has
14441                        // no infinity literal, so the value travels as
14442                        // data (`embryo_core_route_finalize`).
14443                        let mut shell = r.effective_shell(m.experts.len());
14444                        shell.push(f32::NEG_INFINITY);
14445                        meta[base + 30] = pack.put(&shell);
14446                        meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&shared.gate_proj));
14447                        meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&shared.up_proj));
14448                        meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&shared.down_proj));
14449                        for (e, ex) in m.experts.iter().enumerate() {
14450                            meta[base + 32 + e * 3] =
14451                                pack.put(&Self::embryo_qtensor_f32(&ex.gate_proj));
14452                            meta[base + 33 + e * 3] =
14453                                pack.put(&Self::embryo_qtensor_f32(&ex.up_proj));
14454                            meta[base + 34 + e * 3] =
14455                                pack.put(&Self::embryo_qtensor_f32(&ex.down_proj));
14456                        }
14457                    }
14458                    FfnKind::DenseMoe(_) => return None,
14459                }
14460            }
14461            if !full_seen && self.num_layers == 0 {
14462                return None;
14463            }
14464            let id = {
14465                static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
14466                NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
14467            };
14468            let model = crate::gpu::EmbryoGraphModel {
14469                id,
14470                hidden: self.hidden_size,
14471                intermediate: self.intermediate_size,
14472                vocab: self.vocab_size,
14473                layers: self.num_layers,
14474                phase_heads: vmf.map(|c| c.num_heads).unwrap_or(0),
14475                nphase: vmf.map(|c| c.nphase).unwrap_or(0),
14476                phase_dv: vmf.map(|c| c.value_head_dim).unwrap_or(0),
14477                anchor_q_heads: self.num_heads,
14478                anchor_kv_heads: self.num_kv_heads,
14479                anchor_head_dim: self.head_dim,
14480                rotary_dim: self.rotary_dim,
14481                max_seq: self.kv_cache.max_seq_len,
14482                cluster_count,
14483                cluster_size: self.vocab_size / cluster_count,
14484                phase_state_len: vmf.map(|c| c.state_len()).unwrap_or(0),
14485                state_stride,
14486                kv_stride,
14487                norm_gemma: matches!(self.norm_style, NormStyle::Gemma),
14488                phase_mass: vmf.map(|c| c.phase_mass).unwrap_or(0.0),
14489                weights: pack.data,
14490                meta,
14491                lm_head: Self::embryo_qtensor_f32(&self.weights.lm_head),
14492                clusters: clusters.as_ref().clone(),
14493                final_norm: self.weights.final_norm.clone(),
14494                inv_freq: self.inv_freq.as_ref().clone(),
14495                bounded: bounded.is_some(),
14496                kv_layers: if bounded.is_some() {
14497                    bounded_seen
14498                } else {
14499                    self.num_layers
14500                },
14501                state_layers: phase_seen + gdn_seen,
14502                anchor_window,
14503                anchor_sink,
14504                phase_layers: phase_seen,
14505                gdn_layers: gdn_seen,
14506                gdn_heads: gdn.map(|g| g.num_v_heads).unwrap_or(0),
14507                gdn_k_heads: gdn.map(|g| g.num_k_heads).unwrap_or(0),
14508                gdn_dk: gdn.map(|g| g.key_head_dim).unwrap_or(0),
14509                gdn_dv: gdn.map(|g| g.value_head_dim).unwrap_or(0),
14510                gdn_kk: gdn.map(|g| g.conv_kernel).unwrap_or(0),
14511            };
14512            self.embryo_graph = Some(std::sync::Arc::new(model));
14513        }
14514        self.embryo_graph.clone()
14515    }
14516
14517    fn forward_layers_span(
14518        &mut self,
14519        hidden: &[f32],
14520        position: usize,
14521        task_mask: Option<&TaskMask>,
14522        from: usize,
14523        upto: Option<usize>,
14524    ) -> Vec<f32> {
14525        debug_assert!(
14526            from == 0
14527                || (self.dsv4.is_none()
14528                    && self.dsv41.is_none()
14529                    && self.qwen4_exp.is_none()
14530                    && self.g3n.is_none())
14531        );
14532        // Every plain forward — the whole-token Metal graph (`q1_graph_gpu`
14533        // wraps the GDN owners zero-copy and reallocates them on a size
14534        // change) and the CPU layer loop (reads/swaps `linear_state`) —
14535        // must see the previous speculative commit's asynchronous replay
14536        // complete. One mutex probe when nothing is pending.
14537        #[cfg(target_os = "macos")]
14538        if !crate::gpu_metal::wait_replay() {
14539            self.fail_metal_graph("the pending async replay failed before a plain forward");
14540            return vec![0.0; self.hidden_size];
14541        }
14542        if let Some(b) = &mut self.qwen4_exp {
14543            let _ = (task_mask, upto);
14544            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14545            let mut logits = Vec::new();
14546            crate::qwen4_exp::forward_token(
14547                &b.0,
14548                &b.1,
14549                &b.2,
14550                &mut b.3,
14551                token_id,
14552                position,
14553                &self.inv_freq,
14554                self.pool.as_deref(),
14555                &mut logits,
14556                true,
14557            );
14558            self.graph_logits = Some(logits);
14559            return vec![0.0; self.hidden_size];
14560        }
14561        // DeepSeek-V4 runs its own stack: the state is hc_mult copies, and
14562        // the forward returns LOGITS, not a hidden — the head is inside it
14563        // (the final fold sits between the last layer and the norm). The
14564        // token id rides in `hidden[0]`, written by embed_single, because
14565        // the hash layers route by id rather than by content.
14566        if let Some(b) = &mut self.dsv4 {
14567            let _ = (task_mask, upto);
14568            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14569            let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
14570            st.pos = position;
14571            let mut logits = Vec::new();
14572            crate::dsv4::forward_token(
14573                g,
14574                layers,
14575                &cfg,
14576                st,
14577                token_id,
14578                &self.inv_freq,
14579                self.pool.as_deref(),
14580                &mut logits,
14581            );
14582            self.graph_logits = Some(logits);
14583            self.dspark_probe(position, token_id);
14584            // The caller expects a hidden; the logits went out of band, as
14585            // with the fused lm_head path.
14586            return vec![0.0; self.hidden_size];
14587        }
14588        // DeepSeek-V4.1 owns its complete stack and emits logits out of band.
14589        if let Some(b) = &mut self.dsv41 {
14590            let _ = (task_mask, upto);
14591            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14592            let mut logits = Vec::new();
14593            crate::dsv41::forward_token(
14594                &b.0,
14595                &b.1,
14596                &b.2,
14597                &mut b.3,
14598                token_id,
14599                position,
14600                self.pool.as_deref(),
14601                &mut logits,
14602            );
14603            self.graph_logits = Some(logits);
14604            return vec![0.0; self.hidden_size];
14605        }
14606        // Gemma-3n runs its own stack (4 AltUp replicas don't fit this
14607        // loop); `hidden` is the extended embedding from embed_single.
14608        if let Some(b) = &self.g3n {
14609            let _ = (task_mask, upto);
14610            return crate::g3n::g3n_forward(
14611                &b.0,
14612                &b.1,
14613                hidden,
14614                position,
14615                &mut self.kv_cache.layers,
14616                self.num_heads,
14617                self.num_kv_heads,
14618                self.head_dim,
14619                self.pool.as_deref(),
14620            );
14621        }
14622        // Cortiq Embryo owns a separate resident graph: phase recurrent
14623        // state, resonance routing, the GQA anchor KV and hierarchical head
14624        // all execute in one Vulkan submit. It is limited to a complete
14625        // unmasked stack; spans and task masks retain the exact host path.
14626        // `CMF_EMBRYO_DBG=1` names the gate that keeps a token off the
14627        // resident graph — every refusal below is otherwise silent.
14628        if from == 0
14629            && upto.is_none()
14630            && task_mask.is_none()
14631            && self.anchor_core.is_some()
14632            && std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1")
14633        {
14634            static ONCE: std::sync::Once = std::sync::Once::new();
14635            ONCE.call_once(|| {
14636                eprintln!(
14637                    "embryo-dbg: graph_on decode={} prefill={} resident_env={:?} enabled_here={} \
14638                     unsupported={} eligible={}",
14639                    crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode),
14640                    crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill),
14641                    std::env::var("CMF_EMBRYO_RESIDENT").ok(),
14642                    crate::gpu::enabled_here(),
14643                    self.graph_refused(),
14644                    self.embryo_resident_eligible(),
14645                );
14646            });
14647        }
14648        if from == 0
14649            && upto.is_none()
14650            && task_mask.is_none()
14651            // Embryo's recurrent/KV state has no host import path.  Do not
14652            // seed it for a prefill-only graph and then silently decode from
14653            // an empty CPU cache; both phases must select the resident owner.
14654            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14655            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14656            // This whole-token Embryo path remains explicitly opt-in.
14657            // `CMF_GPU_WGPU_GRAPH=1` still enables the mature generic graph,
14658            // but must not silently select this model-specific resident path.
14659            && matches!(
14660                std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14661                Ok("1") | Ok("parallel")
14662            )
14663            && crate::gpu::enabled_here()
14664            && !self.graph_refused()
14665            // Sequence owner: past position zero the device continues only
14666            // a sequence it holds. A host-owned sequence (the graph refused
14667            // at its start, or its prefix was prefilled on the host) keeps
14668            // the host path to its end — never a device attempt at p > 0
14669            // over an empty device image.
14670            && (position == 0 || self.device_sequence_position().is_some())
14671            && self.embryo_resident_eligible()
14672            && let Some(model) = self.ensure_embryo_graph()
14673        {
14674            let mut lg = Vec::new();
14675            if crate::gpu::forward_embryo_graph(&model, self.graph_kv_id, hidden, position, &mut lg)
14676            {
14677                self.graph_logits = Some(lg);
14678                return vec![0.0; self.hidden_size];
14679            }
14680            if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14681                eprintln!("embryo-dbg: forward_embryo_graph refused at position {position}");
14682            }
14683            // The refusal is THIS pipeline's: falling through once is safe
14684            // at position zero (the host owns the sequence from here),
14685            // while a refusal at a later position of a device-owned
14686            // sequence would mix a host KV/state path with a partial
14687            // device sequence.
14688            self.mark_graph_refused();
14689            if position != 0 {
14690                // Fail this sequence alone and leave nothing stale behind:
14691                // no reuse key, no host or device state for the next
14692                // request on this slot to "extend".
14693                self.kv_cache.clear();
14694                self.clear_history();
14695                crate::gpu::graph_kv_reset(self.graph_kv_id);
14696                panic!(
14697                    "resident Embryo graph refused at position {position}; refusing an in-flight host fallback"
14698                );
14699            }
14700        }
14701        let mut h = hidden.to_vec();
14702        // MiMo-V2 expert placement: decided before the graph or the per-op
14703        // arena can claim the budget the expert bank needs.
14704        self.mimo_moe_prepare();
14705        let _mimo_q8 = self.mimo_moe.is_on()
14706            .then(crate::qtensor::enter_full_gpu_q8_scope);
14707        // Split borrows: copy scalars / clone handles so the per-layer
14708        // cfg does not hold `&self` while the KV cache is `&mut`.
14709        let (nh, _nkv, _hd, hs, _rd, eps) = (
14710            self.num_heads,
14711            self.num_kv_heads,
14712            self.head_dim,
14713            self.hidden_size,
14714            self.rotary_dim,
14715            self.rms_eps,
14716        );
14717        let pool = self.pool.clone();
14718        // Opt-in wgpu token-graph attention (discrete Vulkan/DX12): the whole
14719        // attention sub-block runs resident in one submit. Off by default.
14720        // Whole-token wgpu graph: eligibility + arbitration.
14721        //  - explicit CMF_GPU_WGPU_GRAPH forces it on/off;
14722        //  - discrete adapters (4090: decode 76 -> 137 tok/s) and GDN
14723        //    hybrids (recurrent state device-resident, no CPU twin to
14724        //    race) TRUST it;
14725        //  - integrated/mobile adapters RACE it against the normal path
14726        //    at generation granularity (gpu::graph_race_*) — tiled
14727        //    mobile GPUs can turn the ~300-dispatch graph into seconds
14728        //    per token, while a fast phone GPU keeps its win.
14729        let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
14730        let graph_on = match graph_env.as_deref() {
14731            Some("0") => false,
14732            Some("prefill") => false, // decode keeps the per-op path
14733            Some(_) => true,
14734            // Unset: same discrete-only default as every other graph
14735            // site. "Is the GPU on" used to stand in here — which made
14736            // the 0.2 tok/s whole-token graph race-eligible on mobile
14737            // adapters and cost 12-14× on first tokens (cmfmobile
14738            // TUNING.md); integrated GPUs keep the per-op probe path.
14739            None => crate::gpu::wgpu_graph_default(),
14740        };
14741        let graph_trusted =
14742            graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
14743        let race_eligible = graph_on
14744            && upto.is_none()
14745            && task_mask.is_none()
14746            && from == 0
14747            && !self.graph_refused();
14748        let mut tail_start = 0usize;
14749        if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
14750            let t_graph = std::time::Instant::now();
14751            let mut lg = Vec::new();
14752            let mut gl = 0usize;
14753            let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
14754            let declined = built.is_none();
14755            let built = match built {
14756                Some(Ok(hh)) => Some(hh),
14757                Some(Err(())) => {
14758                    // O(1) state was admitted before the device failure; the
14759                    // CPU mirrors are stale by construction.  Clear the whole
14760                    // sequence and stop rather than walking that stale state.
14761                    self.clear_sequence_state();
14762                    self.graph_failed
14763                        .store(true, std::sync::atomic::Ordering::Relaxed);
14764                    self.cancel
14765                        .store(true, std::sync::atomic::Ordering::Relaxed);
14766                    tracing::error!("token graph failed after admission; sequence state cleared");
14767                    return vec![0.0; self.hidden_size];
14768                }
14769                None => None,
14770            };
14771            // Past the transient guards (o1 still collecting, a softcap)
14772            // a refusal is about the weights and will never change —
14773            // remember it instead of walking every layer again next
14774            // token.
14775            if declined && !self.o1_active() && self.attn_softcap == 0.0 {
14776                self.mark_graph_refused();
14777            }
14778            graph_note(built.is_some(), gl, self.num_layers);
14779            if let Some(hh) = built {
14780                let dur = t_graph.elapsed();
14781                if std::env::var("CMF_GRAPH_PROF").is_ok() {
14782                    eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
14783                }
14784                if gl > 0 && gl < self.num_layers {
14785                    // Device prefix: the graph ran layers 0..gl and handed
14786                    // back the boundary hidden — the loop below owns the
14787                    // tail. The prefix layers' KV/state advanced on the
14788                    // device; the tail's advances on the host below. One
14789                    // boundary crossing per token.
14790                    h = hh;
14791                    tail_start = gl;
14792                } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
14793                    if !graph_trusted {
14794                        crate::gpu::graph_race_record(true, dur);
14795                    }
14796                    if !lg.is_empty() {
14797                        // Graph produced logits (final-norm + lm_head folded in) —
14798                        // pad/cap to vocab and hand them to the sampler directly.
14799                        lg.resize(self.vocab_size, 0.0);
14800                        if let Some(c) = self.final_softcap {
14801                            for l in lg.iter_mut() {
14802                                *l = c * (*l / c).tanh();
14803                            }
14804                        }
14805                        self.graph_logits = Some(lg);
14806                    }
14807                    return hh;
14808                }
14809                // Hopeless first graph token: discard it and fall through
14810                // to the normal path. Safe exactly here — the prompt KV is
14811                // still CPU-owned (chunked prefill), so recomputing this
14812                // position is exact; the mirror's extra row is never read
14813                // (the race just settled on the normal path).
14814            }
14815        }
14816        // KIMI-LINEAR HAS NO SPLIT BUG. The 2.6× reported from the
14817        // model rotation (12.2 tok/s on one card against 4.6 on two)
14818        // was a single measurement of a model whose arm arbitration is
14819        // borderline, and it did not survive repetition. Three runs an
14820        // arm, same binary, back to back:
14821        //   probe on : 1 GPU 9.5 / 5.7 / 5.9   2 GPU 7.8 / 13.0 / 13.3
14822        //   pinned   : 1 GPU 5.6 / 5.3 / 5.2   2 GPU 3.5 / 4.2 / 3.4
14823        // With the arms pinned the split costs about 1.45×, which is
14824        // what a layer split costs. With the probe free, TWO CARDS RUN
14825        // FASTER — because for this model the CPU arm wins some op
14826        // classes and the probe finds that.
14827        //
14828        // Two things do stand, and both are measured. The token graph
14829        // builds NOTHING here (`covered 0 of 14 layers [0..14)`), so
14830        // every layer walks per-op on either arm — that is where the
14831        // headroom is, not in the split. And this model's benchmark is
14832        // unusable without `CMF_GPU_PROBE=0`: the arbitration alone
14833        // moves it by more than 2×.
14834        //
14835        // Span runs (network split): the graph covers exactly [from..=upto]
14836        // — one submit per SEGMENT per token. No race: its state is global
14837        // and calibrated on full stacks, so spans take the graph only where
14838        // it is trusted by default (discrete adapters / CMF_GPU_WGPU_GRAPH).
14839        let span = from > 0 || upto.is_some();
14840        if span && graph_on && task_mask.is_none() && graph_trusted {
14841            let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
14842            let mut lg = Vec::new();
14843            let mut gl = 0usize;
14844            let span_res =
14845                self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
14846            let span_res = match span_res {
14847                Some(Ok(hh)) => Some(hh),
14848                Some(Err(())) => {
14849                    self.clear_sequence_state();
14850                    self.graph_failed
14851                        .store(true, std::sync::atomic::Ordering::Relaxed);
14852                    self.cancel
14853                        .store(true, std::sync::atomic::Ordering::Relaxed);
14854                    tracing::error!(
14855                        "span token graph failed after admission; sequence state cleared"
14856                    );
14857                    return vec![0.0; self.hidden_size];
14858                }
14859                None => None,
14860            };
14861            graph_note(span_res.is_some(), gl, upto_excl - from);
14862            if std::env::var("CMF_GPU_DEBUG").is_ok() {
14863                // How much of the span the graph actually covered. A
14864                // prefix of nothing means every layer walks per-op and
14865                // the split's extra cost is elsewhere.
14866                static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
14867                if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
14868                    eprintln!(
14869                        "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
14870                        upto_excl - from,
14871                        span_res.is_some()
14872                    );
14873                }
14874            }
14875            if let Some(hh) = span_res {
14876                if gl == upto_excl - from {
14877                    if !lg.is_empty() {
14878                        lg.resize(self.vocab_size, 0.0);
14879                        if let Some(c) = self.final_softcap {
14880                            for l in lg.iter_mut() {
14881                                *l = c * (*l / c).tanh();
14882                            }
14883                        }
14884                        self.graph_logits = Some(lg);
14885                    }
14886                    crate::gpu::set_layer(-1);
14887                    return hh;
14888                }
14889                // Partial device prefix of the span: CPU owns the tail.
14890                h = hh;
14891                tail_start = from + gl;
14892            }
14893        }
14894        // Layers the host is about to run whose device mirror moved ahead
14895        // of the host cache (a device prefix that shrank since the prompt,
14896        // a batched-prefill prefix longer than this token's): bring their
14897        // rows over first. One comparison per layer when nothing lags.
14898        let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
14899
14900        // A partial graph is an explicit GPU-prefix / CPU-tail split. Keep
14901        // the tail PURE host-side: letting its QTensor hooks re-enter the
14902        // residency arena streams every omitted layer through Vulkan and the
14903        // driver's freed-allocation cache can grow to the full model size
14904        // (25.4 GiB observed with a 14 GiB budget on Granite 30B Q8_2F).
14905        // With a MiMo expert bank the tail is not a whole-layer host
14906        // stream: its experts run from the bank (never the arena) and its
14907        // projections stay per-op on the device, which the bank's placement
14908        // left room for.
14909        let host_tail = tail_start > from;
14910        let _host_tail = (host_tail && !self.mimo_moe.is_on()).then(crate::gpu::enter_cpu_scope);
14911        let automatic_gpu_prefix = self.automatic_gpu_prefix();
14912
14913        let _prof_layers = crate::cpuprof::time(crate::cpuprof::Slot::Layers);
14914        #[cfg(target_os = "macos")]
14915        let mut gpu_skip_until = 0usize;
14916        for li in tail_start.max(from)..self.num_layers {
14917            let _capacity_tail = automatic_gpu_prefix
14918                .filter(|&prefix| li >= prefix && !self.mimo_moe.is_dynamic(li, host_tail))
14919                .map(|_| crate::gpu::enter_cpu_scope());
14920            crate::gpu::set_layer(li as i64); // layer-split GPU/CPU (CMF_GPU_LAYERS)
14921            if let Some(u) = upto {
14922                if li > u {
14923                    break;
14924                }
14925            }
14926            if let Some(mask) = task_mask {
14927                if !mask.layer_alive(li) {
14928                    continue; // dead layer: residual pass-through
14929                }
14930            }
14931            // Whole-block q1 token graph: a run of consecutive q1
14932            // layers — GDN and full attention — executes with one sync
14933            // per CPU attend instead of per op (macOS/Metal).
14934            #[cfg(target_os = "macos")]
14935            {
14936                if li < gpu_skip_until {
14937                    continue;
14938                }
14939                if task_mask.is_none() {
14940                    let end = self.q1_graph_gpu(li, upto, position, &mut h);
14941                    if self
14942                        .graph_failed
14943                        .load(std::sync::atomic::Ordering::Relaxed)
14944                    {
14945                        // The graph may have mutated device state before a
14946                        // command-buffer error. Never continue with a CPU
14947                        // tail or read a stale host mirror after admission.
14948                        return vec![0.0; self.hidden_size];
14949                    }
14950                    if end > li {
14951                        gpu_skip_until = end;
14952                        // Looped Transformer: the graph stopped at a loop
14953                        // boundary — apply final norm before the next iteration.
14954                        if self.is_loop_end(end - 1) && end < self.num_layers {
14955                            h = inference::rms_norm(
14956                                &h,
14957                                &self.weights.final_norm,
14958                                self.rms_eps,
14959                                self.norm_style,
14960                            );
14961                        }
14962                        continue;
14963                    }
14964                }
14965            }
14966
14967            if task_mask.is_none() {
14968                match self.mimo_graph_layer_rows(li, &mut h, &[position]) {
14969                    crate::gpu::BatchGraphOutcome::Completed => continue,
14970                    crate::gpu::BatchGraphOutcome::Failed => return vec![0.0; self.hidden_size],
14971                    crate::gpu::BatchGraphOutcome::Declined => {},
14972                }
14973            }
14974            #[cfg(feature = "gpu")]
14975            self.pull_lagging_host_kv(li, li + 1, position);
14976            let lw = &self.weights.layers[self.phys_layer(li)];
14977            if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
14978                if tp.parse::<usize>().ok() == Some(position) {
14979                    let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
14980                    eprintln!(
14981                        "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
14982                        h[0], h[1]
14983                    );
14984                }
14985            }
14986            // Norm into the pipeline scratch — the returning rms_norm
14987            // allocated twice per layer per token (roadmap §3 P0).
14988            let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
14989            inference::rms_norm_into(
14990                &h,
14991                &lw.input_norm,
14992                self.rms_eps,
14993                self.norm_style,
14994                &mut self.ws.n1,
14995            );
14996            drop(prof);
14997
14998            let attn_out = match &lw.attn {
14999                AttnKind::Mla(w) => {
15000                    let inv_freq_l = self.layer_inv_freq(li);
15001                    let rs = self.layer_rope_scale(li);
15002                    let eps = self.rms_eps;
15003                    let pool = self.pool.clone();
15004                    mla_attention(
15005                        w,
15006                        &self.ws.n1,
15007                        &mut self.kv_cache.layers[li],
15008                        position,
15009                        &inv_freq_l,
15010                        rs,
15011                        eps,
15012                        pool.as_deref(),
15013                    )
15014                }
15015                AttnKind::Linear(w) => {
15016                    let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
15017                    vmf_phase_forward(
15018                        &self.ws.n1,
15019                        w,
15020                        &cfg,
15021                        &mut self.kv_cache.layers[li].linear_state,
15022                        self.pool.as_deref(),
15023                    )
15024                }
15025                AttnKind::Kda(w) => {
15026                    let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
15027                    crate::linear_core::kda_forward(
15028                        &self.ws.n1,
15029                        w,
15030                        &cfg,
15031                        &mut self.kv_cache.layers[li].linear_state,
15032                        self.pool.as_deref(),
15033                    )
15034                }
15035                AttnKind::LinearGdn(w) => {
15036                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
15037                    gdn_forward(
15038                        &self.ws.n1,
15039                        w,
15040                        &cfg,
15041                        &mut self.kv_cache.layers[li].linear_state,
15042                        self.pool.as_deref(),
15043                    )
15044                }
15045                AttnKind::ShortConv(w) => {
15046                    let cfg = self
15047                        .short_conv_cfg
15048                        .expect("short-conv layer without short_conv_cfg");
15049                    short_conv_forward(
15050                        &self.ws.n1,
15051                        w,
15052                        &cfg,
15053                        &mut self.kv_cache.layers[li].linear_state,
15054                        self.pool.as_deref(),
15055                    )
15056                }
15057                AttnKind::Bounded(w) => {
15058                    // Natively bounded anchor: insert into the ring, attend
15059                    // over sinks ∪ window. No position, nothing appended.
15060                    let rope = self
15061                        .bounded_rope
15062                        .clone()
15063                        .expect("bounded layer without an installed rotation table");
15064                    let cfg = crate::bounded::BoundedAttnCfg {
15065                        num_heads: self.num_heads,
15066                        num_kv_heads: self.num_kv_heads,
15067                        head_dim: self.head_dim,
15068                        hidden_size: hs,
15069                        scale: self.attn_scale,
15070                        rope: &rope,
15071                        pool: pool.as_deref(),
15072                    };
15073                    crate::bounded::bounded_attention(
15074                        &self.ws.n1,
15075                        w,
15076                        &mut self.kv_cache.layers[li],
15077                        &cfg,
15078                    )
15079                }
15080                AttnKind::Full {
15081                    wq,
15082                    wk,
15083                    wv,
15084                    wo,
15085                    q_norm,
15086                    k_norm,
15087                    output_gate,
15088                    softplus_gate,
15089                    bias,
15090                } if self.kv_cache.layers[li].o1_sealed() => {
15091                    // O(1) override: decode on the sealed Nyström state
15092                    // instead of the growing KV cache.
15093                    let inv_freq_l = self.layer_inv_freq(li);
15094                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15095                    let cfg = QwenAttnCfg {
15096                        num_heads: self.layer_num_heads(li),
15097                        num_kv_heads: nkv_l,
15098                        head_dim: hd_l,
15099                        hidden_size: hs,
15100                        position,
15101                        inv_freq: &inv_freq_l,
15102                        rotary_dim: rd_l,
15103                        scale: self.attn_scale,
15104                        softcap: self.attn_softcap,
15105                        window: None,
15106                        v_norm: self.attn_v_norm,
15107                        qk_norm_after_rope: self.qk_norm_after_rope,
15108                        gate_sigmoid: self.proj_gate_sigmoid,
15109                        q_norm: q_norm.as_deref(),
15110                        k_norm: k_norm.as_deref(),
15111                        output_gate: *output_gate,
15112                        softplus_gate: softplus_gate
15113                            .as_ref()
15114                            .map(|(gate, per_head)| (gate, *per_head)),
15115                        rope_scale: self.layer_rope_scale(li),
15116                        bias: bias
15117                            .as_ref()
15118                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
15119                        rms_eps: eps,
15120                        norm_style: self.norm_style,
15121                        pool: pool.as_deref(),
15122                        v_head_dim: self.layer_v_dim(li),
15123                    };
15124                    attention::qwen_attention_nystrom(
15125                        &self.ws.n1,
15126                        wq,
15127                        wk,
15128                        wv,
15129                        wo,
15130                        &mut self.kv_cache.layers[li],
15131                        &cfg,
15132                    )
15133                }
15134                AttnKind::Full {
15135                    wq,
15136                    wk,
15137                    wv,
15138                    wo,
15139                    q_norm,
15140                    k_norm,
15141                    output_gate,
15142                    softplus_gate,
15143                    bias,
15144                } => 'attn: {
15145                    // wgpu token-graph attention (opt-in): whole sub-block in
15146                    // one submit, device K/V mirror. q1 only, no gate/bias/mask.
15147                    // Its kernel has no window, sink or narrow-V slot and
15148                    // one mirror geometry: such models stay on the CPU attend.
15149                    let dropin_reason =
15150                        graph_on.then(|| self.graph_attn_decline_reason()).flatten();
15151                    if let Some(reason) = dropin_reason {
15152                        self.note_graph_decline("wgpu attn dropin", reason);
15153                    }
15154                    if graph_on
15155                        && dropin_reason.is_none()
15156                        && !*output_gate
15157                        && softplus_gate.is_none()
15158                        && self.attention_heads_per_layer.is_none()
15159                        && bias.is_none()
15160                        && task_mask.is_none()
15161                    {
15162                        let inv_freq_l = self.layer_inv_freq(li);
15163                        let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15164                        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
15165                        if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
15166                            wq.mapped_q1(),
15167                            wk.mapped_q1(),
15168                            wv.mapped_q1(),
15169                            wo.mapped_q1(),
15170                        ) {
15171                            let gm = gm.clone();
15172                            let mut out = vec![0f32; hs];
15173                            let cache = &self.kv_cache.layers[li];
15174                            if crate::gpu::attn_dropin(
15175                                &gm,
15176                                self.graph_kv_id,
15177                                li,
15178                                &self.ws.n1,
15179                                qi,
15180                                ki,
15181                                vi,
15182                                oi,
15183                                q_norm.as_deref(),
15184                                k_norm.as_deref(),
15185                                self.qk_norm_after_rope,
15186                                &inv_freq_l,
15187                                nh,
15188                                nkv_l,
15189                                hd_l,
15190                                rd_l,
15191                                hs,
15192                                position,
15193                                self.kv_cache.max_seq_len,
15194                                gemma,
15195                                eps as f32,
15196                                cache.k_heads(),
15197                                cache.v_heads(),
15198                                &mut out,
15199                            ) {
15200                                break 'attn out;
15201                            }
15202                        }
15203                    }
15204                    let masked = task_mask
15205                        .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
15206                        .unwrap_or(false);
15207                    let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
15208                    // The masked kernel knows one pipeline-wide geometry and
15209                    // RoPE table, no window and no sink.
15210                    let plain = self.layer_attn_plain(li);
15211                    match (masked, f32_view) {
15212                        // Historical masked path (f32 slices; the loader
15213                        // keeps masked models in f32).
15214                        (true, (Some(q), Some(k), Some(v), Some(o))) if plain => {
15215                            let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
15216                            attention::multi_head_attention(
15217                                &self.ws.n1,
15218                                q,
15219                                k,
15220                                v,
15221                                o,
15222                                &mut self.kv_cache.layers[li],
15223                                self.num_heads,
15224                                self.num_kv_heads,
15225                                self.head_dim,
15226                                self.hidden_size,
15227                                position,
15228                                &active_heads,
15229                                &self.inv_freq,
15230                            )
15231                        }
15232                        (masked, _) => {
15233                            if masked {
15234                                tracing::warn!(
15235                                    "layer {li}: head mask on quantized weights or on a \
15236                                     window/sink/per-layer-geometry layer not supported \
15237                                     yet — executing dense"
15238                                );
15239                            }
15240                            let inv_freq_l = self.layer_inv_freq(li);
15241                            let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15242                            let cfg = QwenAttnCfg {
15243                                num_heads: self.layer_num_heads(li),
15244                                num_kv_heads: nkv_l,
15245                                head_dim: hd_l,
15246                                hidden_size: hs,
15247                                position,
15248                                inv_freq: &inv_freq_l,
15249                                rotary_dim: rd_l,
15250                                scale: self.attn_scale,
15251                                softcap: self.attn_softcap,
15252                                window: self.layer_window(li),
15253                                v_norm: self.attn_v_norm,
15254                                qk_norm_after_rope: self.qk_norm_after_rope,
15255                                gate_sigmoid: self.proj_gate_sigmoid,
15256                                q_norm: q_norm.as_deref(),
15257                                k_norm: k_norm.as_deref(),
15258                                output_gate: *output_gate,
15259                                softplus_gate: softplus_gate
15260                                    .as_ref()
15261                                    .map(|(gate, per_head)| (gate, *per_head)),
15262                                rope_scale: self.layer_rope_scale(li),
15263                                bias: bias
15264                                    .as_ref()
15265                                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
15266                                rms_eps: eps,
15267                                norm_style: self.norm_style,
15268                                pool: pool.as_deref(),
15269                                v_head_dim: self.layer_v_dim(li),
15270                            };
15271                            attention::qwen_attention(
15272                                &self.ws.n1,
15273                                wq,
15274                                wk,
15275                                wv,
15276                                wo,
15277                                &mut self.kv_cache.layers[li],
15278                                &cfg,
15279                            )
15280                        }
15281                    }
15282                }
15283            };
15284            // Gemma sandwich norm: normalize the attention branch before
15285            // it joins the residual stream.
15286            let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
15287                Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
15288                None => attn_out,
15289            };
15290            let lw = &self.weights.layers[self.phys_layer(li)];
15291            let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
15292            inference::add_rmsnorm_fused_into(
15293                &mut h,
15294                &attn_out,
15295                &lw.post_norm,
15296                self.rms_eps,
15297                self.norm_style,
15298                &mut self.ws.p1,
15299            );
15300            drop(prof);
15301            let mut attn_out = attn_out;
15302            attention::recycle_buf(&mut attn_out);
15303            let post_normed = &self.ws.p1;
15304
15305            let ffn_masked = task_mask
15306                .map(|m| m.ffn_active_count(li) < self.intermediate_size)
15307                .unwrap_or(false);
15308            // One masked dense CONTRACT, dispatched by cost. The
15309            // activation-zeroing arm (the batched sweep's, validated
15310            // against the replica to 0.8%) computes the FULL fused FFN
15311            // and zeroes the dead — right whenever most neurons live.
15312            // The sparse arm reads ONLY active rows and down columns —
15313            // per-row dots are slower per element than the fused kernel,
15314            // so it pays only once the mask is deep enough. The 0.5
15315            // crossover is first-principles (fused kernels run ~2x the
15316            // per-row dot throughput); a shallow specialist (95% alive)
15317            // stays fused, a --target-sparsity bake flips arms on its
15318            // own weight.
15319            let ffn_out = match (ffn_masked, &lw.ffn) {
15320                // A defragged tube layer answers its own mask: the core
15321                // always runs, each tube runs when its bit is on, and
15322                // the tubes that are off are never read from the mmap.
15323                (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
15324                    let row = task_mask
15325                        .and_then(|tm| tm.ffn_masks.get(li))
15326                        .map(|v| v.as_slice());
15327                    tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
15328                }
15329                (true, FfnKind::Dense(d)) => {
15330                    let tm = task_mask.unwrap();
15331                    let alive = tm.ffn_active_count(li);
15332                    let deep = alive * 2 <= self.intermediate_size;
15333                    if deep && d.down_proj.sparse_col_ok() && !d.gate_proj.has_prism_contract() {
15334                        let active = tm.ffn_active_indices(li);
15335                        sparse_ffn_quant(
15336                            d,
15337                            post_normed,
15338                            &active,
15339                            self.hidden_size,
15340                            self.pool.as_deref(),
15341                        )
15342                    } else if deep
15343                        && let (Some(g), Some(u), Some(dn)) = (
15344                            d.gate_proj.as_f32(),
15345                            d.up_proj.as_f32(),
15346                            d.down_proj.as_f32(),
15347                        )
15348                    {
15349                        let active = tm.ffn_active_indices(li);
15350                        inference::sparse_ffn_forward(
15351                            post_normed,
15352                            g,
15353                            u,
15354                            dn,
15355                            self.hidden_size,
15356                            self.intermediate_size,
15357                            &active,
15358                            self.pool.as_deref(),
15359                        )
15360                    } else {
15361                        let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
15362                        dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
15363                    }
15364                }
15365                (true, FfnKind::Moe(m)) => {
15366                    // MoE is sparse by expert selection; a task mask
15367                    // narrows the ROUTABLE set via its expert fields
15368                    // (spec §5) when it carries them.
15369                    let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
15370                    ffn_forward(
15371                        &lw.ffn,
15372                        post_normed,
15373                        self.pool.as_deref(),
15374                        allowed.as_deref(),
15375                    )
15376                }
15377                (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
15378                    dm,
15379                    post_normed,
15380                    &h,
15381                    self.rms_eps,
15382                    self.norm_style,
15383                    self.pool.as_deref(),
15384                ),
15385                (false, _) => match &lw.ffn {
15386                    FfnKind::DenseMoe(dm) => dense_moe_ffn(
15387                        dm,
15388                        post_normed,
15389                        &h,
15390                        self.rms_eps,
15391                        self.norm_style,
15392                        self.pool.as_deref(),
15393                    ),
15394                    FfnKind::Moe(m)
15395                        if task_mask.is_none() && self.mimo_moe.is_dynamic(li, host_tail) =>
15396                    {
15397                        moe_ffn_banked(&mut self.mimo_moe, li, m, post_normed, self.pool.as_deref())
15398                    }
15399                    _ => {
15400                        let allowed = match (&lw.ffn, task_mask) {
15401                            (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
15402                            _ => None,
15403                        };
15404                        ffn_forward(
15405                            &lw.ffn,
15406                            post_normed,
15407                            self.pool.as_deref(),
15408                            allowed.as_deref(),
15409                        )
15410                    }
15411                },
15412            };
15413            let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
15414                Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
15415                None => ffn_out,
15416            };
15417            for (i, &f) in ffn_out.iter().enumerate() {
15418                h[i] += f;
15419            }
15420            let mut ffn_out = ffn_out;
15421            attention::recycle_buf(&mut ffn_out);
15422
15423            // Gemma-4: the layer output is scaled by a learned scalar.
15424            if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
15425                for v in h.iter_mut() {
15426                    *v *= sc;
15427                }
15428            }
15429            // CMF_LAYER_DUMP: this position's hidden after layer li.
15430            if self.layer_dump.is_some() {
15431                self.dump_layer_row(position, li, &h);
15432            }
15433
15434            // Looped Transformer: apply final norm at the end of each loop iteration.
15435            // Nanbeige 4.2: after layer 21 (virtual), apply norm before looping back to layer 0.
15436            if self.is_loop_end(li) && li + 1 < self.num_layers {
15437                h = inference::rms_norm(
15438                    &h,
15439                    &self.weights.final_norm,
15440                    self.rms_eps,
15441                    self.norm_style,
15442                );
15443            }
15444
15445            // Dynamic routing φ capture (on-policy): the
15446            // EMA of the post-residual hidden at the router's phi_layer,
15447            // updated as the context evolves during decode.
15448            if self.dyn_phi_layer == Some(li) {
15449                self.update_dyn_phi(&h);
15450            }
15451        }
15452        crate::gpu::set_layer(-1); // layers done — lm_head outside layer-split
15453        if let Some(t) = t_race_cpu {
15454            crate::gpu::graph_race_record(false, t.elapsed());
15455        }
15456
15457        h
15458    }
15459
15460    /// EMA of φ at the router layer (rolling, weight 0.2 = ~5-token
15461    /// horizon). First observation seeds it exactly.
15462    fn update_dyn_phi(&mut self, h: &[f32]) {
15463        const A: f32 = 0.2;
15464        if self.dyn_phi_ema.len() != h.len() {
15465            self.dyn_phi_ema = vec![0.0; h.len()];
15466            self.dyn_phi_seen = 0;
15467        }
15468        if self.dyn_phi_seen == 0 {
15469            self.dyn_phi_ema.copy_from_slice(h);
15470        } else {
15471            for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
15472                *e = (1.0 - A) * *e + A * v;
15473            }
15474        }
15475        self.dyn_phi_seen += 1;
15476    }
15477
15478    /// Current router φ (EMA at phi_layer); empty until first capture.
15479    pub fn dyn_phi(&self) -> &[f32] {
15480        &self.dyn_phi_ema
15481    }
15482
15483    /// Enable/disable φ capture at the router layer, reset the EMA.
15484    pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
15485        self.dyn_phi_layer = layer;
15486        self.dyn_phi_ema.clear();
15487        self.dyn_phi_seen = 0;
15488    }
15489
15490    /// Skills eligible for dynamic switching: (index, id, phi_layer).
15491    pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
15492        let Some(model) = &self.model else {
15493            return Vec::new();
15494        };
15495        model
15496            .header
15497            .skills
15498            .iter()
15499            .enumerate()
15500            .filter_map(|(i, sk)| {
15501                // A v2 record routes only through the request-level
15502                // backbone-gated decision (its status and gate are
15503                // checked there), never per token.
15504                if sk.is_v2() {
15505                    return None;
15506                }
15507                let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
15508                let sel = sk.selection.as_ref()?;
15509                (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
15510            })
15511            .collect()
15512    }
15513
15514    /// Index of the currently overlaid skill (None = backbone).
15515    pub fn active_skill(&self) -> Option<usize> {
15516        self.dyn_active
15517    }
15518
15519    /// Enable dynamic per-token skill routing: build the hysteresis
15520    /// router from the container's routable skills, start φ capture at
15521    /// their (shared) phi_layer. Returns the number of routable skills
15522    /// (0 = nothing to route; router stays off). Idempotent.
15523    pub fn enable_dynamic_routing(&mut self) -> usize {
15524        use crate::swarm::{DynRouter, RoutableSkill};
15525        let Some(model) = self.model.clone() else {
15526            return 0;
15527        };
15528        // Router policy v2 (spec §9.4) routes per REQUEST: the backbone is
15529        // the default and only the backbone-gated decision may pick a
15530        // skill. A per-token switch would bypass that gate (and change
15531        // the O(1) state mid-sequence), so the hysteresis router never
15532        // runs on such a file; the caller keeps the request-level
15533        // decision.
15534        if let Some(r) = &model.header.router {
15535            tracing::warn!(
15536                "dynamic routing disabled: this file declares router policy '{}' with \
15537                 granularity \"{}\" — the request-level decision applies instead",
15538                r.policy,
15539                r.granularity
15540            );
15541            return 0;
15542        }
15543        // Format-v2 skill records (bit SKILLS_V2) without a router policy:
15544        // their status/gate contract ("auto-routing requires active +
15545        // measured") lives in the backbone-gated decision only — the
15546        // hysteresis router would switch into a quarantined record
15547        // (fail-open). Refuse the whole file, not just its v2 records.
15548        if model.required_features & cortiq_core::format::features::SKILLS_V2 != 0
15549            || model.header.skills.iter().any(|s| s.is_v2())
15550        {
15551            tracing::warn!(
15552                "dynamic routing disabled: this file carries format-v2 skill records \
15553                 (SKILLS_V2) — they route per request through a router policy only"
15554            );
15555            return 0;
15556        }
15557        // A blend materialized f32 working tensors into the layers; there
15558        // is no single skill index to revert from → refuse (honest).
15559        if self.dyn_blend_loaded {
15560            tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
15561            return 0;
15562        }
15563        // A statically-overlaid skill that is NOT FFN-eligible can't be
15564        // cheaply reverted at generation start → refuse rather than
15565        // silently keep it overlaid.
15566        if let Some(a) = self.dyn_active {
15567            if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
15568                tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
15569                return 0;
15570            }
15571        }
15572        let hidden = self.hidden_size;
15573        let mut skills = Vec::new();
15574        for (idx, id, _phi) in self.dynamic_skills() {
15575            if let Some(sel) = model.header.skills[idx].selection.as_ref() {
15576                if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
15577                    skills.push(rs);
15578                }
15579            }
15580        }
15581        if skills.is_empty() {
15582            return 0;
15583        }
15584        // Skills should share a phi_layer; warn (not fail) if they don't.
15585        let phi = skills[0].phi_layer;
15586        if skills.iter().any(|s| s.phi_layer != phi) {
15587            tracing::warn!("routable skills disagree on phi_layer; using {phi}");
15588        }
15589        let n = skills.len();
15590        self.set_dyn_phi_layer(Some(phi));
15591        self.dyn_router = Some(DynRouter::new(skills));
15592        n
15593    }
15594
15595    /// Human-readable switch log from the last dynamic-routed generation.
15596    pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
15597        self.dyn_router
15598            .as_ref()
15599            .map(|r| r.switches.clone())
15600            .unwrap_or_default()
15601    }
15602
15603    /// LM head: hidden → logits [vocab_size]. The dominant matvec of
15604    /// every decode step — row-parallel on the worker pool.
15605    fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
15606        let _mimo_q8 = self.mimo_moe.is_on()
15607            .then(crate::qtensor::enter_full_gpu_q8_scope);
15608        let rows = self.weights.lm_head.rows();
15609        let mut logits = attention::take_buf(rows.min(self.vocab_size));
15610        // Banked MiMo uses the same exact projection family for the
15611        // plain/draft head and the batched verification head. Read both
15612        // scale planes in-place instead of preparing per-op scale buffers.
15613        let served = self.mimo_moe.is_on() && crate::gpu::mimo_q8_short_enabled()
15614            && rows == self.vocab_size && !self.weights.lm_head.has_prism_contract()
15615            && self.weights.lm_head.graph_weight().is_some_and(|(model, idx, kind, _)| {
15616                kind == 7 && crate::gpu::q82_short_rows(model, idx, hidden, 1,
15617                    rows, self.hidden_size, &mut logits)
15618            });
15619        if !served {
15620            self.weights.lm_head.matvec(hidden, &mut logits, self.pool.as_deref());
15621        }
15622        logits.resize(self.vocab_size, 0.0);
15623        if let Some(m) = self.logit_multiplier {
15624            for l in logits.iter_mut() {
15625                *l *= m;
15626            }
15627        }
15628        if let Some(c) = self.final_softcap {
15629            for l in logits.iter_mut() {
15630                *l = c * (*l / c).tanh();
15631            }
15632        }
15633        if let Some(cm) = self.head_clusters.as_ref() {
15634            self.hierarchical_head_logprobs(hidden, cm, &mut logits);
15635        }
15636        logits
15637    }
15638
15639    /// Two-level head (Cortiq Embryo): in place, logits[v] ← log p(v) =
15640    /// (lc[c] − lse(lc)) + (logit[v] − lse over v's cluster block), c = v / S.
15641    fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
15642        let h = hidden.len();
15643        let ncl = cm.len() / h.max(1);
15644        if ncl == 0 || logits.len() % ncl != 0 {
15645            return;
15646        }
15647        let cs = logits.len() / ncl;
15648        // cluster logits + log-softmax
15649        let mut lc = vec![0.0f32; ncl];
15650        for c in 0..ncl {
15651            let row = &cm[c * h..(c + 1) * h];
15652            let mut s = 0.0f32;
15653            for j in 0..h {
15654                s += row[j] * hidden[j];
15655            }
15656            lc[c] = s;
15657        }
15658        let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15659        let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
15660        for c in 0..ncl {
15661            let blk = &mut logits[c * cs..(c + 1) * cs];
15662            let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15663            let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
15664            let add = lc[c] - lse - bl;
15665            for v in blk.iter_mut() {
15666                *v += add;
15667            }
15668        }
15669    }
15670
15671    /// Prefill `ids` and return the next-token logits — what the model
15672    /// would predict next, WITHOUT committing to generation (introspection
15673    /// for `cortiq explain`). Clears and repopulates the KV cache; leaves
15674    /// the active overlay untouched.
15675    pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
15676        #[cfg(target_os = "macos")]
15677        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
15678        self.clear_sequence_state();
15679        // This helper is used by the pooled classification endpoint, where
15680        // every request is a fresh sequence. The shared reset also clears the
15681        // wgpu token graph's device-side recurrent state.
15682        crate::gpu::graph_race_begin_generation();
15683        if task_mask.is_none() {
15684            self.o1_begin();
15685        }
15686        let mut hidden = vec![0.0f32; self.hidden_size];
15687        for (pos, &id) in ids.iter().enumerate() {
15688            let emb = self.embed_single(id);
15689            hidden = self.forward_layers(&emb, pos, task_mask);
15690        }
15691        if let Err(err) = self.o1_seal_checked() {
15692            self.o1_fail(err);
15693        }
15694        // Stacks that own their head (V4, V4.1, Qwen3.8-Flash-Next, GLM-5)
15695        // return a zero hidden and hand the logits out of band.
15696        if let Some(logits) = self.graph_logits.take() {
15697            return logits;
15698        }
15699        inference::rms_norm_into(
15700            &hidden,
15701            &self.weights.final_norm,
15702            self.rms_eps,
15703            self.norm_style,
15704            &mut self.ws.n1,
15705        );
15706        self.lm_head_forward(&self.ws.n1)
15707    }
15708}
15709
15710/// Convenience: deterministic tiny pipeline for tests.
15711pub fn create_test_pipeline(
15712    hidden_size: usize,
15713    intermediate_size: usize,
15714    num_heads: usize,
15715    num_kv_heads: usize,
15716    head_dim: usize,
15717    num_layers: usize,
15718    vocab_size: usize,
15719) -> Pipeline {
15720    // Small pseudo-random weights: constant weights make attention
15721    // degenerate and hide indexing bugs.
15722    let synth = |n: usize, salt: usize| -> Vec<f32> {
15723        (0..n)
15724            .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
15725            .collect()
15726    };
15727    let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
15728        QTensor::from_f32(synth(rows * cols, salt), rows, cols)
15729    };
15730    let layer_weights: Vec<LayerWeights> = (0..num_layers)
15731        .map(|li| LayerWeights {
15732            input_norm: vec![1.0; hidden_size],
15733            post_norm: vec![1.0; hidden_size],
15734            attn_out_norm: None,
15735            ffn_out_norm: None,
15736            layer_scale: None,
15737            ffn: FfnKind::Dense(DenseFfn {
15738                gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
15739                up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
15740                down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
15741                act: Act::Silu,
15742                down_t: None,
15743                segs: Vec::new(),
15744            }),
15745            attn: AttnKind::Full {
15746                bias: None,
15747                wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
15748                wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
15749                wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
15750                wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
15751                q_norm: None,
15752                k_norm: None,
15753                output_gate: false,
15754                softplus_gate: None,
15755            },
15756        })
15757        .collect();
15758
15759    Pipeline::new(
15760        Tokenizer::byte_level(),
15761        PipelineWeights {
15762            embed_tokens: qt(vocab_size, hidden_size, 100),
15763            layers: layer_weights,
15764            lm_head: qt(vocab_size, hidden_size, 200),
15765            final_norm: vec![1.0; hidden_size],
15766        },
15767        hidden_size,
15768        intermediate_size,
15769        num_heads,
15770        num_kv_heads,
15771        head_dim,
15772        num_layers,
15773        num_layers, // physical_layers = num_layers (non-looped)
15774        false,      // loop_final_norm
15775        vocab_size,
15776        1e-6,
15777        10_000.0,
15778        NormStyle::Qwen,
15779        4096,
15780        SamplerConfig {
15781            seed: Some(42),
15782            ..Default::default()
15783        },
15784    )
15785}
15786
15787/// Batched dense-FFN: gate/up/down via matmat (element-wise the same
15788/// math as b × dense_ffn — the same dot kernels).
15789/// One mask bit, LSB-first per byte — `TaskMask::ffn_active_indices`'s
15790/// convention.
15791#[inline]
15792fn mask_bit(row: &[u8], j: usize) -> bool {
15793    (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
15794}
15795
15796/// Zero the CLOSED neurons' activations in a [rows × inter] panel — the
15797/// masked-inference fast path's whole trick: full fused quant compute,
15798/// then the mask lands on the ACTIVATIONS, which is arithmetically the
15799/// pruned network without touching a quantized weight byte. Whole open
15800/// bytes (0xFF = 8 open neurons) skip in one test.
15801/// `CMF_FFN_MASK_GAIN` — Patent 12 FIG. 4, variance-preserving
15802/// rescaling: truncation removes a share of the layer's output energy,
15803/// so the survivors are scaled up to put the variance back where the
15804/// downstream norm expects it. A scalar here; per layer it is
15805/// `sqrt(total energy / kept energy)`.
15806fn mask_gain() -> f32 {
15807    static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
15808    *G.get_or_init(|| {
15809        std::env::var("CMF_FFN_MASK_GAIN")
15810            .ok()
15811            .and_then(|v| v.parse().ok())
15812            .unwrap_or(1.0)
15813    })
15814}
15815
15816fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
15817    // With CMF_FFN_MEANFILL a closed neuron contributes its average
15818    // instead of nothing — same bytes read, one constant restored.
15819    let fill = meanfill().and_then(|(i, v)| {
15820        let li = crate::gpu::cur_layer();
15821        (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
15822    });
15823    for r in 0..rows {
15824        let base = r * inter;
15825        for (bi, &byte) in row.iter().enumerate() {
15826            if byte == 0xFF {
15827                continue;
15828            }
15829            let j0 = bi * 8;
15830            for bit in 0..8 {
15831                let j = j0 + bit;
15832                if j < inter && byte & (1 << bit) == 0 {
15833                    g[base + j] = fill.map_or(0.0, |f| f[j]);
15834                }
15835            }
15836        }
15837    }
15838    let gain = mask_gain();
15839    if gain != 1.0 {
15840        for v in g[..rows * inter].iter_mut() {
15841            *v *= gain;
15842        }
15843    }
15844}
15845
15846/// True when neuron `i`'s bit is set (no mask = everything runs).
15847#[inline]
15848fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
15849    row.is_none_or(|r| mask_bit(r, i))
15850}
15851
15852/// Every bit below `n` set — the common case for a tube file's CORE,
15853/// where only the tube bits vary per task.
15854fn all_bits_on(row: &[u8], n: usize) -> bool {
15855    (0..n).all(|i| mask_bit(row, i))
15856}
15857
15858/// `CMF_TUBE_TOPK` — how many tubes a TOKEN may open (0 = the task mask
15859/// decides alone). This is the dense FFN read as a mixture: the tubes
15860/// are the experts a k-means over `gate_proj` rows found, and the token
15861/// picks among them. `CMF_TUBE_SCORE=gate` scores a tube by its own
15862/// gate (realizable: only `up`/`down` of the losers go unread),
15863/// `=oracle` scores by the true `silu(gate)·up` mass (the ceiling —
15864/// only `down` is saved, and the selection has read what it predicts).
15865fn tube_topk() -> usize {
15866    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
15867    *K.get_or_init(|| {
15868        std::env::var("CMF_TUBE_TOPK")
15869            .ok()
15870            .and_then(|v| v.parse().ok())
15871            .unwrap_or(0)
15872    })
15873}
15874
15875fn tube_score_oracle() -> bool {
15876    static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15877    *O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
15878}
15879
15880/// The routed arm of `tube_ffn`: a token opens only its best `k` tubes.
15881/// At `b == 1` (decode) the losers are genuinely never read — that is
15882/// the speed. At `b > 1` (the scoring sweep) every tube is computed and
15883/// the losers' activations are zeroed instead: same arithmetic, so the
15884/// perplexity is the routed model's, measured without a per-token
15885/// gather in the middle of a GEMM.
15886fn tube_ffn_routed(
15887    d: &DenseFfn,
15888    xs: &[f32],
15889    b: usize,
15890    pool: Option<&Pool>,
15891    mask_row: Option<&[u8]>,
15892    k: usize,
15893) -> Vec<f32> {
15894    let hidden = d.down_proj.rows();
15895    let core = d.gate_proj.rows();
15896    let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
15897    let mut out = match (b, core_full, mask_row) {
15898        (1, true, _) => dense_ffn(d, xs, pool),
15899        (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
15900        (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
15901        (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
15902    };
15903    let cand: Vec<usize> = (0..d.segs.len())
15904        .filter(|&i| tube_bit(mask_row, d.segs[i].start))
15905        .collect();
15906    if cand.is_empty() {
15907        return out;
15908    }
15909    // gate (and, where the score or the batch needs it, up) per tube.
15910    // The SCORE is taken at the point the serving path could take it:
15911    // off the gate alone, or off the finished activation for the oracle.
15912    let oracle = tube_score_oracle();
15913    let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
15914    let mut scores = vec![0f32; b * cand.len()];
15915    for (ci, &i) in cand.iter().enumerate() {
15916        let seg = &d.segs[i];
15917        let w = seg.width;
15918        let mut g = vec![0.0f32; b * w];
15919        if b == 1 {
15920            seg.gate.matvec(xs, &mut g, pool);
15921        } else {
15922            seg.gate.matmat(xs, b, &mut g, pool);
15923        }
15924        for v in g.iter_mut() {
15925            *v = Act::Silu.combine(*v, 1.0);
15926        }
15927        if !oracle {
15928            for t in 0..b {
15929                scores[t * cand.len() + ci] =
15930                    g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15931            }
15932        }
15933        if oracle || b > 1 {
15934            let mut u = vec![0.0f32; b * w];
15935            if b == 1 {
15936                seg.up.matvec(xs, &mut u, pool);
15937            } else {
15938                seg.up.matmat(xs, b, &mut u, pool);
15939            }
15940            for (a, &v) in g.iter_mut().zip(u.iter()) {
15941                *a *= v;
15942            }
15943            if oracle {
15944                for t in 0..b {
15945                    scores[t * cand.len() + ci] =
15946                        g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15947                }
15948            }
15949        }
15950        acts.push(g);
15951    }
15952    // per-token scores and the winners
15953    let keep = k.min(cand.len());
15954    let mut scratch: Vec<f32> = Vec::new();
15955    for t in 0..b {
15956        let mut sc: Vec<(f32, usize)> = (0..cand.len())
15957            .map(|ci| (scores[t * cand.len() + ci], ci))
15958            .collect();
15959        sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
15960        let mut alive = vec![false; cand.len()];
15961        for &(_, ci) in sc.iter().take(keep) {
15962            alive[ci] = true;
15963        }
15964        if b > 1 {
15965            for (ci, a) in acts.iter_mut().enumerate() {
15966                if !alive[ci] {
15967                    let w = d.segs[cand[ci]].width;
15968                    a[t * w..(t + 1) * w].fill(0.0);
15969                }
15970            }
15971        } else {
15972            // decode: finish only the winners — the losers' up/down
15973            // (and, with the gate score, everything but their gate)
15974            // are never touched.
15975            for (ci, &i) in cand.iter().enumerate() {
15976                if !alive[ci] {
15977                    continue;
15978                }
15979                let seg = &d.segs[i];
15980                let w = seg.width;
15981                let g = &mut acts[ci];
15982                if !tube_score_oracle() {
15983                    scratch.clear();
15984                    scratch.resize(w, 0.0);
15985                    seg.up.matvec(xs, &mut scratch, pool);
15986                    for (a, &v) in g.iter_mut().zip(scratch.iter()) {
15987                        *a *= v;
15988                    }
15989                }
15990                let mut acc = vec![0.0f32; hidden];
15991                seg.down.matvec(g, &mut acc, pool);
15992                for (o, a) in out.iter_mut().zip(&acc) {
15993                    *o += *a;
15994                }
15995            }
15996        }
15997    }
15998    if b > 1 {
15999        for (ci, &i) in cand.iter().enumerate() {
16000            let seg = &d.segs[i];
16001            let mut acc = vec![0.0f32; b * hidden];
16002            seg.down.matmat(&acts[ci], b, &mut acc, pool);
16003            for (o, a) in out.iter_mut().zip(&acc) {
16004                *o += *a;
16005            }
16006        }
16007    }
16008    out
16009}
16010
16011/// FFN of a defragged tube layer: the always-on core plus the tubes the
16012/// task mask switches on. Each tube is a normal tensor triple, so the
16013/// same kernels run it and an inactive tube's bytes are never read —
16014/// that is the whole point of the defrag (a scattered mask cannot skip
16015/// bytes; a contiguous one is just a smaller matrix).
16016fn tube_ffn(
16017    d: &DenseFfn,
16018    xs: &[f32],
16019    b: usize,
16020    pool: Option<&Pool>,
16021    mask_row: Option<&[u8]>,
16022) -> Vec<f32> {
16023    if tube_topk() > 0 {
16024        return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
16025    }
16026    let hidden = d.down_proj.rows();
16027    let core = d.gate_proj.rows();
16028    let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
16029    let mut out = match (b, core_full, mask_row) {
16030        (1, true, _) => dense_ffn(d, xs, pool),
16031        (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
16032        (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
16033        (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
16034    };
16035    TUBE_SCRATCH.with(|sc| {
16036        let mut sc = sc.borrow_mut();
16037        let [g, u, acc] = &mut *sc;
16038        for seg in &d.segs {
16039            if !tube_bit(mask_row, seg.start) {
16040                continue;
16041            }
16042            let w = seg.width;
16043            g.resize(b * w, 0.0);
16044            if b == 1
16045                && d.act == Act::Silu
16046                && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
16047            {
16048                // g holds silu(gate)·up.
16049            } else {
16050                u.resize(b * w, 0.0);
16051                if b == 1 {
16052                    QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
16053                } else {
16054                    seg.gate.matmat(xs, b, g, pool);
16055                    seg.up.matmat(xs, b, u, pool);
16056                }
16057                for i in 0..b * w {
16058                    g[i] = d.act.combine(g[i], u[i]);
16059                }
16060            }
16061            acc.resize(b * hidden, 0.0);
16062            acc.fill(0.0);
16063            if b == 1 {
16064                seg.down.matvec(g, acc, pool);
16065            } else {
16066                seg.down.matmat(g, b, acc, pool);
16067            }
16068            for (o, a) in out.iter_mut().zip(acc.iter()) {
16069                *o += *a;
16070            }
16071        }
16072        out
16073    })
16074}
16075
16076thread_local! {
16077    /// gate / up / down-accumulator scratch for the tube loop — a tube
16078    /// runs once per layer per token, and a fresh Vec each time is a
16079    /// malloc per tube per layer per token.
16080    static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
16081        const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
16082}
16083
16084/// CMF_PREFILL_PROF: cumulative ns of the batched walk's attention and
16085/// FFN halves (all layers, all chunks).
16086static PREFILL_SPLIT: [std::sync::atomic::AtomicU64; 2] =
16087    [std::sync::atomic::AtomicU64::new(0), std::sync::atomic::AtomicU64::new(0)];
16088
16089fn prefill_prof_on() -> bool {
16090    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16091    *ON.get_or_init(|| std::env::var_os("CMF_PREFILL_PROF").is_some())
16092}
16093
16094fn dense_ffn_batch(
16095    d: &DenseFfn,
16096    xs: &[f32],
16097    b: usize,
16098    pool: Option<&Pool>,
16099    mask_row: Option<&[u8]>,
16100) -> Vec<f32> {
16101    let inter = d.gate_proj.rows();
16102    let hidden = d.down_proj.rows();
16103    // Fused on-device SwiGLU when the device is in play: three separate
16104    // `matmat` calls are three round trips per layer, and the gate/up
16105    // panels (b × inter — 22 MB each at a 512-token chunk) cross the bus
16106    // twice for nothing. The kernel already existed for the image DiT;
16107    // the LLM prefill was simply never wired to it. A task mask needs the
16108    // activations on the host between the halves, so it keeps the CPU
16109    // arm below.
16110    // SiLU and the exact GELU (Spark-X2.5) have device arms; `q4_ffn_act`
16111    // hands SiLU to the very entry points this used to call.
16112    let fused_act = d.act.graph_act();
16113    if mask_row.is_none()
16114        && fused_act.is_some()
16115        && b >= 32
16116        && crate::gpu::enabled_here()
16117        && !crate::gpu::mm_killed()
16118        // The refit pass needs this layer's activations on the host; the
16119        // fused chain keeps them on the device. Refusing it here costs
16120        // one round trip and keeps every GEMM on the card — the
16121        // alternative was running the whole calibration on the CPU.
16122        && refit_dir().is_none()
16123        // Same for the mass/hit probes. The accumulator at the bottom of
16124        // this function only sees `g` when `g` came back to the host, so
16125        // a fused batch would leave it summing nothing — a probe that
16126        // reports zeros rather than failing, which is worse.
16127        && !ffn_probe_active()
16128    {
16129        if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16130            d.gate_proj.mapped_q4t(),
16131            d.up_proj.mapped_q4t(),
16132            d.down_proj.mapped_q4t(),
16133        ) {
16134            let mut out = vec![0.0f32; b * hidden];
16135            let act = fused_act.expect("checked above");
16136            if crate::gpu::q4_ffn_act(model, w1, w3, w2, xs, b, hidden, inter, false, act, &mut out)
16137            {
16138                return out;
16139            }
16140        }
16141        // The q4tp twin (same kernel family, scale from the row ladder) —
16142        // the DiT has run it in production since the pipeline containers;
16143        // the LLM prefill was simply never wired to it, so a q4tp model's
16144        // prefill panels stayed on the CPU.
16145        if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16146            d.gate_proj.mapped_q4tp(),
16147            d.up_proj.mapped_q4tp(),
16148            d.down_proj.mapped_q4tp(),
16149        ) {
16150            let mut out = vec![0.0f32; b * hidden];
16151            let act = fused_act.expect("checked above");
16152            if crate::gpu::q4_ffn_act(model, w1, w3, w2, xs, b, hidden, inter, true, act, &mut out)
16153            {
16154                return out;
16155            }
16156        }
16157        // Every other device codec — int8, or gate/up and down in different
16158        // codecs: the three GEMMs with the panels kept on the card. Without
16159        // it a q8_2f prefill read both b·inter panels home, folded them on
16160        // one host thread, and sent the result back for down.
16161        // CMF_FFN_KEEP=0 keeps the per-GEMM path (A/B).
16162        if let (true, Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16163            std::env::var("CMF_FFN_KEEP").as_deref() != Ok("0"),
16164            d.gate_proj.mapped_device_gemm(),
16165            d.up_proj.mapped_device_gemm(),
16166            d.down_proj.mapped_device_gemm(),
16167        ) {
16168            let mut out = vec![0.0f32; b * hidden];
16169            let act = fused_act.expect("checked above");
16170            if crate::gpu::ffn_act_keep(model, w1, w3, w2, xs, b, hidden, inter, act, &mut out) {
16171                return out;
16172            }
16173        }
16174    }
16175    let mut g = vec![0.0f32; b * inter];
16176    d.gate_proj.matmat(xs, b, &mut g, pool);
16177    let mut u = vec![0.0f32; b * inter];
16178    d.up_proj.matmat(xs, b, &mut u, pool);
16179    if gate_topk() > 0 && d.act == Act::Silu {
16180        for t in 0..b {
16181            let row = &mut g[t * inter..(t + 1) * inter];
16182            for v in row.iter_mut() {
16183                *v = Act::Silu.combine(*v, 1.0);
16184            }
16185            keep_top_k(row, gate_topk());
16186        }
16187        for i in 0..b * inter {
16188            g[i] *= u[i];
16189        }
16190    } else {
16191        for i in 0..b * inter {
16192            g[i] = d.act.combine(g[i], u[i]);
16193        }
16194    }
16195    if let Some(row) = mask_row {
16196        zero_masked_cols(&mut g, b, inter, row);
16197    }
16198    if oracle_topk() > 0 {
16199        for t in 0..b {
16200            keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
16201        }
16202    }
16203    let mut out = vec![0.0f32; b * hidden];
16204    d.down_proj.matmat(&g, b, &mut out, pool);
16205    if refit_dir().is_some() {
16206        let li = crate::gpu::cur_layer();
16207        if li >= 0 {
16208            refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
16209        }
16210    }
16211    // The DTG-MA probe, on the batched path: one prefill sweep gives the
16212    // same per-neuron statistic the per-position probe does, and on a 27B
16213    // that is minutes instead of hours.
16214    FFN_PROBE.with(|pr| {
16215        if let Some(acc) = pr.borrow_mut().as_mut() {
16216            let li = crate::gpu::cur_layer();
16217            if li < 0 {
16218                return;
16219            }
16220            let Some(row) = acc.get_mut(li as usize) else {
16221                return;
16222            };
16223            let sq = probe_sq();
16224            for t in 0..b {
16225                for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
16226                    *a += if sq {
16227                        (v as f64) * (v as f64)
16228                    } else {
16229                        (v as f64).abs()
16230                    };
16231                }
16232            }
16233        }
16234    });
16235    out
16236}
16237
16238/// Batched MoE-FFN: router batched, positions are GROUPED by expert —
16239/// an expert's weights are read once for all its positions in the chunk
16240/// (the main prefill-GEMM win on MoE: 960MB/token of 35B experts).
16241/// Accumulate per-channel activation energy for `CMF_RMS_TRACE`.
16242fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
16243    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16244    static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16245    let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
16246    let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
16247    if (!on && !dump) || b == 0 {
16248        return;
16249    }
16250    let hidden = xs.len() / b;
16251    if on {
16252        let mut acc = m.act_sq.borrow_mut();
16253        if acc.len() < hidden {
16254            acc.resize(hidden, 0.0);
16255        }
16256        for t in 0..b {
16257            let row = &xs[t * hidden..(t + 1) * hidden];
16258            for (a, &v) in acc.iter_mut().zip(row) {
16259                *a += (v as f64) * (v as f64);
16260            }
16261        }
16262    }
16263    if dump {
16264        // Cap the capture: the covariance needs a few thousand rows, and a
16265        // whole prefill of every layer would be gigabytes for no extra rank.
16266        let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
16267            .ok()
16268            .and_then(|v| v.parse().ok())
16269            .unwrap_or(4096);
16270        let mut rows = m.act_rows.borrow_mut();
16271        if rows.len() < cap * hidden {
16272            let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
16273            rows.extend_from_slice(&xs[..take * hidden]);
16274        }
16275    }
16276}
16277
16278/// Send-able cursor over a Vec-of-Vecs: each pool worker writes only its
16279/// own slots (disjoint by construction in the caller).
16280#[derive(Clone, Copy)]
16281struct SendVecs(*mut Vec<f32>);
16282unsafe impl Send for SendVecs {}
16283unsafe impl Sync for SendVecs {}
16284impl SendVecs {
16285    #[inline]
16286    fn at(self, i: usize) -> *mut Vec<f32> {
16287        unsafe { self.0.add(i) }
16288    }
16289}
16290
16291fn moe_ffn_batch(
16292    m: &MoeFfn,
16293    xs: &[f32],
16294    b: usize,
16295    hidden: usize,
16296    pool: Option<&Pool>,
16297    allowed: Option<&[bool]>,
16298) -> Vec<f32> {
16299    accumulate_act(m, xs, b);
16300    let ne = m.experts.len();
16301    let mut logits = vec![0.0f32; b * ne];
16302    match &m.resonance {
16303        Some(r) => {
16304            let hdim = xs.len() / b.max(1);
16305            for bi in 0..b {
16306                r.scores(
16307                    &xs[bi * hdim..(bi + 1) * hdim],
16308                    &mut logits[bi * ne..(bi + 1) * ne],
16309                );
16310            }
16311        }
16312        None => m.router.matmat(xs, b, &mut logits, pool),
16313    }
16314
16315    // Assignments: expert → [(position, weight)] — same routing as
16316    // moe_ffn, per position (see `moe_route`).
16317    let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
16318    {
16319        let mut st = m.stats.borrow_mut();
16320        if st.len() < ne {
16321            st.resize(ne, 0);
16322        }
16323        for bi in 0..b {
16324            let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
16325            for &e in &idx {
16326                st[e] += 1;
16327                assign[e].push((bi, p[e] / wsum));
16328            }
16329        }
16330    }
16331
16332    let mut out = vec![0.0f32; b * hidden];
16333    let cols = m.experts[0].gate_proj.cols();
16334    let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
16335        let sb = list.len();
16336        let mut sub = vec![0.0f32; sb * cols];
16337        for (k, &(bi, _)) in list.iter().enumerate() {
16338            sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16339        }
16340        let eo = dense_ffn_batch(d, &sub, sb, pool, None);
16341        for (k, &(bi, w)) in list.iter().enumerate() {
16342            for i in 0..hidden {
16343                out[bi * hidden + i] += w * eo[k * hidden + i];
16344            }
16345        }
16346    };
16347    // Routed experts: the panels are TINY (b·top_k spread over every
16348    // expert — a few positions each), so a pool dispatch per expert is
16349    // pure barrier cost. Invert the parallelism: workers take WHOLE
16350    // experts (serial math inside), then one deterministic scatter in
16351    // expert order — the exact accumulation order the serial loop had.
16352    let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
16353    if pool.is_some() && active.len() >= 8 {
16354        let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
16355        {
16356            let panel_ptr = SendVecs(panels.as_mut_ptr());
16357            // Capture only the expert table: `m` itself carries RefCell
16358            // stats and must not cross the pool boundary.
16359            let experts = &m.experts;
16360            let (active_r, assign_r) = (&active, &assign);
16361            let inherit_cpu = crate::gpu::inherit_cpu_scope();
16362            let run = |start: usize, end: usize| {
16363                let _cpu_scope = inherit_cpu();
16364                for ai in start..end {
16365                    let e = active_r[ai];
16366                    let list = &assign_r[e];
16367                    let sb = list.len();
16368                    let mut sub = vec![0.0f32; sb * cols];
16369                    for (k, &(bi, _)) in list.iter().enumerate() {
16370                        sub[k * cols..(k + 1) * cols]
16371                            .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16372                    }
16373                    // SAFETY: each worker owns a disjoint panels[ai].
16374                    unsafe {
16375                        *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
16376                    }
16377                }
16378            };
16379            match pool {
16380                Some(p) => p.run_rows(active.len(), &run),
16381                None => run(0, active.len()),
16382            }
16383        }
16384        for (ai, &e) in active.iter().enumerate() {
16385            for (k, &(bi, w)) in assign[e].iter().enumerate() {
16386                let eo = &panels[ai][k * hidden..(k + 1) * hidden];
16387                for i in 0..hidden {
16388                    out[bi * hidden + i] += w * eo[i];
16389                }
16390            }
16391        }
16392    } else {
16393        for &e in &active {
16394            run_expert(&m.experts[e], &assign[e], &mut out);
16395        }
16396    }
16397    if let Some((se, gate)) = &m.shared {
16398        let all: Vec<(usize, f32)> = if let Some(gate) = gate {
16399            let mut gl = vec![0.0f32; b];
16400            gate.matmat(xs, b, &mut gl, pool);
16401            (0..b)
16402                .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
16403                .collect()
16404        } else {
16405            (0..b).map(|bi| (bi, 1.0)).collect()
16406        };
16407        run_expert(se, &all, &mut out);
16408    }
16409    out
16410}
16411
16412/// Decode-exact multi-token MoE — the MiMo speculative verify's FFN. Row
16413/// `r` of the result is bit-identical to `moe_ffn(m, x_r)` on the CPU
16414/// (`moe_ffn_cpu` → `moe_ffn_cpu_batched`): router matvec per row, the same
16415/// routing, the same int8 gate/up/SiLU and down terms
16416/// (`QTensor::moe_gate_up_rows` / `moe_down_rows`), and the row's experts
16417/// summed in ITS route order from 0. What the rows share is the weight
16418/// traffic: each routed expert is read once for every row that picked it.
16419/// (`moe_ffn_batch`, the prompt path, groups the same way but sums in
16420/// expert-index order and runs blocked kernels on wide groups — close, not
16421/// bit-equal to decode.) Any layer the kernels do not cover, or a device
16422/// that could answer `moe_ffn` itself, walks `moe_ffn` row by row.
16423fn moe_ffn_rows_exact(
16424    m: &MoeFfn,
16425    xs: &[f32],
16426    b: usize,
16427    hidden: usize,
16428    pool: Option<&Pool>,
16429) -> Vec<f32> {
16430    let mut out = vec![0.0f32; b * hidden];
16431    let per_row = |out: &mut [f32]| {
16432        for r in 0..b {
16433            let o = moe_ffn(m, &xs[r * hidden..(r + 1) * hidden], pool, None);
16434            out[r * hidden..(r + 1) * hidden].copy_from_slice(&o);
16435        }
16436    };
16437    let covered = !crate::gpu::enabled_here()
16438        && moe_batch_enabled()
16439        && m.shared.is_none()
16440        && m.resonance.is_none()
16441        && FFN_PROBE.with(|pr| pr.borrow().is_none())
16442        && m.experts.iter().all(|d| d.act == Act::Silu);
16443    if !covered {
16444        per_row(&mut out);
16445        return out;
16446    }
16447    let ne = m.experts.len();
16448    // Routing, row by row, exactly as `moe_ffn`.
16449    let mut routes: Vec<(Vec<usize>, Vec<f32>)> = Vec::with_capacity(b);
16450    for r in 0..b {
16451        let x = &xs[r * hidden..(r + 1) * hidden];
16452        accumulate_act(m, x, 1);
16453        let mut logits = vec![0.0f32; ne];
16454        m.router.matvec(x, &mut logits, pool);
16455        let (idx, p, wsum) = moe_route(&logits, m, None);
16456        {
16457            let mut st = m.stats.borrow_mut();
16458            if st.len() < ne {
16459                st.resize(ne, 0);
16460            }
16461            for &e in &idx {
16462                st[e] += 1;
16463            }
16464        }
16465        let w: Vec<f32> = idx
16466            .iter()
16467            .map(|&e| p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]))
16468            .collect();
16469        routes.push((idx, w));
16470    }
16471    if routes.iter().any(|(idx, _)| idx.is_empty()) {
16472        per_row(&mut out);
16473        return out;
16474    }
16475    // Group the (row, expert) picks by expert, in first-seen order.
16476    let mut experts: Vec<usize> = Vec::new();
16477    let mut groups: Vec<Vec<usize>> = Vec::new();
16478    for (r, (idx, _)) in routes.iter().enumerate() {
16479        for &e in idx {
16480            match experts.iter().position(|&x| x == e) {
16481                Some(g) => groups[g].push(r),
16482                None => {
16483                    experts.push(e);
16484                    groups.push(vec![r]);
16485                }
16486            }
16487        }
16488    }
16489    let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
16490    let inter = m.experts[experts[0]].gate_proj.rows();
16491    let pairs: Vec<(&QTensor, &QTensor)> = experts
16492        .iter()
16493        .map(|&e| (&m.experts[e].gate_proj, &m.experts[e].up_proj))
16494        .collect();
16495    let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
16496    if !QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut gs, pool) {
16497        per_row(&mut out);
16498        return out;
16499    }
16500    let downs: Vec<&QTensor> = experts.iter().map(|&e| &m.experts[e].down_proj).collect();
16501    let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
16502    let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; hidden]).collect();
16503    if !QTensor::moe_down_rows(&downs, &lens, &gs, &mut ds, pool) {
16504        per_row(&mut out);
16505        return out;
16506    }
16507    // Where each (row, expert) term landed in the flat pair list.
16508    let mut slot = std::collections::HashMap::with_capacity(n_pairs);
16509    let mut p = 0usize;
16510    for (g, &e) in experts.iter().enumerate() {
16511        for &r in &groups[g] {
16512            slot.insert((r, e), p);
16513            p += 1;
16514        }
16515    }
16516    for (r, (idx, w)) in routes.iter().enumerate() {
16517        let terms: Vec<(&[f32], f32)> = idx
16518            .iter()
16519            .zip(w)
16520            .map(|(&e, &we)| (ds[slot[&(r, e)]].as_slice(), we))
16521            .collect();
16522        let row = &mut out[r * hidden..(r + 1) * hidden];
16523        for (i, dst) in row.iter_mut().enumerate() {
16524            // `moe_down_many`'s per-row sum: from 0, in route order.
16525            let mut acc = 0f32;
16526            for (d, we) in &terms {
16527                acc += we * d[i];
16528            }
16529            *dst = acc;
16530        }
16531    }
16532    out
16533}
16534
16535thread_local! {
16536    /// gate/up activation scratch for the dense FFN paths (single uses
16537    /// two slots, the fused pair all four) — these were fresh
16538    /// intermediate-size Vecs on every layer of every token.
16539    static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
16540        const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
16541}
16542
16543/// Dense SwiGLU FFN through QTensor matvecs (any storage).
16544fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16545    // Per-token sparsity, when the file was built for it: gate first,
16546    // then only the chosen neurons' up/down rows leave the mmap.
16547    if gate_topk() > 0
16548        && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
16549    {
16550        return out;
16551    }
16552    // Whole-FFN GPU submit (этап 4.2 increment): gate → silu·up → down
16553    // chained in ONE command buffer with the intermediate activations
16554    // resident on the device — 3 per-op polls become 1 per layer. The
16555    // moe_block backend already implements exactly this chain; a dense
16556    // FFN is one expert with weight 1. Runtime probe: the chain still
16557    // pays one submit+poll per layer — alternate it against the pure-CPU
16558    // FFN and keep whichever is faster on this machine.
16559    // q1 FFNs offload at any practical size: the q1 CPU kernel is
16560    // compute-bound, so the UMA threshold logic does not apply — the
16561    // probe measures and decides either way.
16562    // The fused GPU block has no descriptor-aware Prism path: it would either
16563    // consume an unrotated activation or decline after inspecting the mixed
16564    // q2tp/q4tp tensors.  Do not let that structural refusal enter the FFN
16565    // probe's CPU_ONLY scope; the ordinary body below dispatches each matrix
16566    // through QTensor::matvec, which owns the signed FWHT + affine q2tp route.
16567    let prism_body = d.gate_proj.has_prism_contract()
16568        || d.up_proj.has_prism_contract()
16569        || d.down_proj.has_prism_contract();
16570    if !prism_body
16571        && crate::gpu::enabled_here()
16572        && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
16573    {
16574        let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
16575            crate::gpu::ProbeArm::Gpu
16576        } else {
16577            crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
16578        };
16579        match arm {
16580            crate::gpu::ProbeArm::Gpu => {
16581                let t0 = std::time::Instant::now();
16582                if let Some(out) = dense_ffn_gpu(d, x, pool) {
16583                    crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
16584                    return out;
16585                }
16586                // Declined: no timing exists, so say so. Silence here is
16587                // what left `ffn` undecided for 9000 calls and cost a
16588                // failed device attempt on half of them.
16589                crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
16590            }
16591            crate::gpu::ProbeArm::CpuTimed => {
16592                let t0 = std::time::Instant::now();
16593                let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16594                crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
16595                return out;
16596            }
16597            crate::gpu::ProbeArm::Cpu => {
16598                return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16599            }
16600        }
16601    }
16602    dense_ffn_cpu(d, x, pool)
16603}
16604
16605/// The pure-CPU dense-FFN body (also the fallback of every GPU refusal).
16606fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16607    let inter = d.gate_proj.rows();
16608    FFN_SCRATCH.with(|s| {
16609        let mut s = s.borrow_mut();
16610        let [g, u, ..] = &mut *s;
16611        g.resize(inter, 0.0);
16612        // Fused gate+up+silu: one dispatch, no separate silu pass.
16613        // Falls back to matvec_many + silu loop for unsupported dtypes.
16614        if gate_topk() > 0 {
16615            // Gate first, select, and only then pay for `up`: the
16616            // measurement arm computes both and zeroes the losers, which
16617            // is the same arithmetic.
16618            u.resize(inter, 0.0);
16619            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16620            for i in 0..inter {
16621                g[i] = Act::Silu.combine(g[i], 1.0);
16622            }
16623            keep_top_k(g, gate_topk());
16624            for i in 0..inter {
16625                g[i] *= u[i];
16626            }
16627        } else if d.act == Act::Silu && {
16628            let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16629            QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
16630        } {
16631            // g now holds silu(gate)·up directly.
16632        } else {
16633            u.resize(inter, 0.0);
16634            // Multi-matrix job: gate+up under one pool dispatch.
16635            let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16636            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16637            for i in 0..inter {
16638                g[i] = d.act.combine(g[i], u[i]);
16639            }
16640        }
16641        // DTG-MA bake probe (Patent 2): accumulate this layer's
16642        // per-neuron activation mass while a probe pass is active.
16643        // `CMF_FFN_PROBE_TOPK=k` switches the statistic from mass to a
16644        // HIT COUNT — how many tokens rank the neuron in their own top
16645        // k. Mass asks "how loud is this neuron overall", the count
16646        // asks "how often does this task actually need it", and the two
16647        // rank neurons differently whenever a few tokens are loud.
16648        FFN_PROBE.with(|pr| {
16649            if let Some(acc) = pr.borrow_mut().as_mut() {
16650                let li = crate::gpu::cur_layer();
16651                if li >= 0 {
16652                    if let Some(row) = acc.get_mut(li as usize) {
16653                        match probe_topk() {
16654                            0 if probe_sq() => {
16655                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16656                                    *a += (v as f64) * (v as f64);
16657                                }
16658                            }
16659                            0 if probe_signed() => {
16660                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16661                                    *a += v as f64;
16662                                }
16663                            }
16664                            0 => {
16665                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16666                                    *a += (v as f64).abs();
16667                                }
16668                            }
16669                            k => {
16670                                let n = g.len();
16671                                let k = k.min(n);
16672                                let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16673                                let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16674                                    b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16675                                });
16676                                let thr = *kth;
16677                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16678                                    if v.abs() >= thr {
16679                                        *a += 1.0;
16680                                    }
16681                                }
16682                            }
16683                        }
16684                    }
16685                }
16686            }
16687        });
16688        if oracle_topk() > 0 {
16689            keep_top_k(g, oracle_topk());
16690        }
16691        {
16692            let li = crate::gpu::cur_layer();
16693            if li >= 0 {
16694                adump_row(li as usize, g);
16695            }
16696        }
16697        let mut out = attention::take_buf(d.down_proj.rows());
16698        let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnDown);
16699        d.down_proj.matvec(g, &mut out, pool);
16700        out
16701    })
16702}
16703
16704/// Online accumulators for the AWNP refit of a narrowed FFN.
16705///
16706/// The refit needs `Gss = A_SᵀA_S` and `YA = YᵀA_S` per layer, where `A_S`
16707/// are the calibration activations of the KEPT neurons and `Y` the full
16708/// FFN output. Both are small enough to hold; the thing that is not is
16709/// the activations they are built from — a 27B layer would dump a
16710/// gigabyte per thousand tokens. So they are accumulated as the
16711/// calibration runs and written once at the end.
16712///
16713/// `CMF_FFN_REFIT=<dir>` holds `support.<L>.u32` (a u32 count then the
16714/// kept indices) for every layer to accumulate; `CMF_FFN_REFIT_FROM/TO`
16715/// bound the layer span so the accumulators fit in RAM.
16716pub struct RefitAcc {
16717    pub support: Vec<u32>,
16718    pub gss: Vec<f32>,
16719    pub ya: Vec<f32>,
16720    pub hidden: usize,
16721    pub tokens: u64,
16722    /// Activations staged transposed ([ns, t] and [hidden, t]) until the
16723    /// batch is worth a GEMM. The product costs `ns²` to move and add
16724    /// REGARDLESS of how many tokens went into it, so folding 16 chunks
16725    /// into one call cuts that cost 16× — it was 15 TB of traffic per
16726    /// calibration pass at one call per 256 tokens.
16727    pub buf_g: Vec<f32>,
16728    pub buf_o: Vec<f32>,
16729    pub buf_t: usize,
16730}
16731
16732/// The product buffer is SHARED across layers — one 473 MB allocation,
16733/// not one per layer (that was 30 GB of nothing on a 64-layer model).
16734/// It lives under the same lock as the accumulators.
16735type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
16736
16737static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
16738    std::sync::OnceLock::new();
16739
16740/// Is an FFN probe accumulator installed on this thread? The fused GPU
16741/// FFN must decline while one is, or the probe silently measures zero.
16742fn ffn_probe_active() -> bool {
16743    FFN_PROBE.with(|p| p.borrow().is_some())
16744}
16745
16746fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
16747    REFIT
16748        .get_or_init(|| {
16749            std::env::var("CMF_FFN_REFIT").ok().map(|d| {
16750                (
16751                    d,
16752                    std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
16753                )
16754            })
16755        })
16756        .as_ref()
16757}
16758
16759/// Accumulate one prefill panel into the layer's refit statistics.
16760fn refit_accumulate(
16761    li: usize,
16762    g: &[f32],
16763    b: usize,
16764    inter: usize,
16765    out: &[f32],
16766    hidden: usize,
16767    pool: Option<&Pool>,
16768) {
16769    let Some((dir, map)) = refit_dir() else {
16770        return;
16771    };
16772    static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16773    let (from, to) = *SPAN.get_or_init(|| {
16774        let g = |k: &str, d: usize| {
16775            std::env::var(k)
16776                .ok()
16777                .and_then(|v| v.parse().ok())
16778                .unwrap_or(d)
16779        };
16780        (
16781            g("CMF_FFN_REFIT_FROM", 0),
16782            g("CMF_FFN_REFIT_TO", usize::MAX),
16783        )
16784    });
16785    if li < from || li > to {
16786        return;
16787    }
16788    let mut guard = map.lock().unwrap();
16789    let (map, shared) = &mut *guard;
16790    let acc = match map.entry(li) {
16791        std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
16792        std::collections::hash_map::Entry::Vacant(e) => {
16793            let path = format!("{dir}/support.{li}.u32");
16794            let Ok(bytes) = std::fs::read(&path) else {
16795                eprintln!("refit: no {path} — layer {li} skipped");
16796                return;
16797            };
16798            let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
16799            let support: Vec<u32> = bytes[4..4 + n * 4]
16800                .chunks_exact(4)
16801                .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16802                .collect();
16803            eprintln!(
16804                "refit: layer {li} support {n} ({:.0} MB of accumulator)",
16805                (n * n + hidden * n) as f64 * 4.0 / 1e6
16806            );
16807            e.insert(RefitAcc {
16808                gss: vec![0.0; n * n],
16809                ya: vec![0.0; hidden * n],
16810                buf_g: Vec::new(),
16811                buf_o: Vec::new(),
16812                buf_t: 0,
16813                support,
16814                hidden,
16815                tokens: 0,
16816            })
16817        }
16818    };
16819    let ns = acc.support.len();
16820    // Stage this chunk transposed; the GEMM fires once the batch is full.
16821    let cap = refit_batch();
16822    if acc.buf_g.is_empty() {
16823        acc.buf_g = vec![0.0; ns * cap];
16824        acc.buf_o = vec![0.0; hidden * cap];
16825    }
16826    let take = b.min(cap - acc.buf_t);
16827    for t in 0..take {
16828        let col = acc.buf_t + t;
16829        for (j, &n) in acc.support.iter().enumerate() {
16830            acc.buf_g[j * cap + col] = g[t * inter + n as usize];
16831        }
16832        for h in 0..hidden {
16833            acc.buf_o[h * cap + col] = out[t * hidden + h];
16834        }
16835    }
16836    acc.buf_t += take;
16837    acc.tokens += take as u64;
16838    if acc.buf_t < cap {
16839        return;
16840    }
16841    let bt = acc.buf_t;
16842    acc.buf_t = 0;
16843    // The GEMM WRITES its C (it zeroes the accumulators it uses), so the
16844    // chunk product lands in scratch and is added on — the one thing that
16845    // silently turns a Gram over 13 000 tokens into a Gram over 256.
16846    // Both products are `C[n, m] += X[n, b] · Yᵀ[b, m]` with X and Y
16847    // stored row-major [·, b] — exactly `gemm_nt_f32`'s shape, so the
16848    // card does them when it is up (this is the whole calibration's
16849    // cost: O(|S|²) per token, 2.9 PFLOP for a 27B pass). The tiled CPU
16850    // loop stays as the fallback. Neither accumulates, so the product
16851    // lands in scratch and is added on.
16852    let RefitAcc {
16853        gss,
16854        ya,
16855        buf_g,
16856        buf_o,
16857        ..
16858    } = acc;
16859    let need = (ns * ns).max(hidden * ns);
16860    if shared.len() < need {
16861        shared.resize(need, 0.0);
16862    }
16863    let scratch = &mut shared[..];
16864    let _ = bt;
16865    if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
16866        add_into(gss, &scratch[..ns * ns], pool);
16867        if crate::gpu::gemm_nt_f32_transient(
16868            buf_o,
16869            buf_g,
16870            &mut scratch[..hidden * ns],
16871            hidden,
16872            cap,
16873            ns,
16874        ) {
16875            add_into(ya, &scratch[..hidden * ns], pool);
16876        } else {
16877            accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16878        }
16879    } else {
16880        accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
16881        accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16882    }
16883    // No zeroing: the batch is always filled exactly (cap is a multiple
16884    // of the prefill chunk), and a memset of 178 MB a layer would cost
16885    // more than the GEMM.
16886}
16887
16888/// `CMF_FFN_REFIT_BATCH` — tokens staged before each GEMM (default 4096).
16889fn refit_batch() -> usize {
16890    static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16891    *B.get_or_init(|| {
16892        std::env::var("CMF_FFN_REFIT_BATCH")
16893            .ok()
16894            .and_then(|v| v.parse().ok())
16895            .unwrap_or(4096)
16896    })
16897}
16898
16899/// `c[m, n] += Σ_t left[m, t]·right[n, t]` — both operands transposed,
16900/// the CPU fallback for the staged batch.
16901fn accum_outer_t(
16902    c: &mut [f32],
16903    m: usize,
16904    n: usize,
16905    b: usize,
16906    left: &[f32],
16907    right: &[f32],
16908    pool: Option<&Pool>,
16909) {
16910    let ptr = SendMut(c.as_mut_ptr());
16911    let body = |i: usize| {
16912        let ptr = &ptr;
16913        let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
16914        for t in 0..b {
16915            let a = left[i * b + t];
16916            if a == 0.0 {
16917                continue;
16918            }
16919            for (j, o) in row.iter_mut().enumerate() {
16920                *o += a * right[j * b + t];
16921            }
16922        }
16923    };
16924    match pool {
16925        Some(p) if m > 1 => p.run_rows(m, &|s, e| {
16926            for i in s..e {
16927                body(i);
16928            }
16929        }),
16930        _ => {
16931            for i in 0..m {
16932                body(i);
16933            }
16934        }
16935    }
16936}
16937
16938/// `dst += src`, spread over the pool — at 118 M floats a layer this is
16939/// not a loop to leave on one core.
16940fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
16941    let n = dst.len().min(src.len());
16942    match pool {
16943        Some(p) if n >= 1 << 16 => {
16944            let ptr = SendMut(dst.as_mut_ptr());
16945            let f = |s: usize, e: usize| {
16946                let ptr = &ptr;
16947                for blk in s..e {
16948                    let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
16949                    for i in a..b {
16950                        unsafe { *ptr.0.add(i) += src[i] };
16951                    }
16952                }
16953            };
16954            p.run_rows(n.div_ceil(4096), &f);
16955        }
16956        _ => {
16957            for (d, v) in dst.iter_mut().zip(&src[..n]) {
16958                *d += *v;
16959            }
16960        }
16961    }
16962}
16963
16964/// `c[m, n] += Σ_t left[t, m]·right[t, n]`, with `left` stored [m, t] and
16965/// `right` [t, n]. Tiled over the rows of `c` so a tile stays in cache
16966/// while each token's `right` row streams past it once, and parallel
16967/// over tiles.
16968fn accum_outer(
16969    c: &mut [f32],
16970    m: usize,
16971    n: usize,
16972    b: usize,
16973    left: &[f32],
16974    right: &[f32],
16975    pool: Option<&Pool>,
16976) {
16977    const TILE: usize = 32;
16978    let tiles = m.div_ceil(TILE);
16979    let cp = SendMut(c.as_mut_ptr());
16980    let body = |ti: usize| {
16981        let cp = &cp;
16982        let i0 = ti * TILE;
16983        let i1 = (i0 + TILE).min(m);
16984        for t in 0..b {
16985            let r = &right[t * n..t * n + n];
16986            for i in i0..i1 {
16987                let a = left[i * b + t];
16988                if a == 0.0 {
16989                    continue;
16990                }
16991                // SAFETY: tiles partition c's rows; workers never overlap.
16992                let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
16993                for (o, v) in row.iter_mut().zip(r) {
16994                    *o += a * *v;
16995                }
16996            }
16997        }
16998    };
16999    match pool {
17000        Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
17001            for ti in s..e {
17002                body(ti);
17003            }
17004        }),
17005        _ => {
17006            for ti in 0..tiles {
17007                body(ti);
17008            }
17009        }
17010    }
17011}
17012
17013/// Write what the calibration accumulated: `gss.<L>.f32` and `ya.<L>.f32`.
17014pub fn refit_flush() -> usize {
17015    let Some((dir, map)) = refit_dir() else {
17016        return 0;
17017    };
17018    let guard = map.lock().unwrap();
17019    let mut n = 0;
17020    for (li, acc) in guard.0.iter() {
17021        // A silently truncated write here is a Gram that reshapes to
17022        // nothing an hour later — say it out loud instead.
17023        let w = |name: &str, v: &[f32]| {
17024            let path = format!("{dir}/{name}.{li}.f32");
17025            let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
17026            match std::fs::write(&path, &bytes) {
17027                Ok(()) => {}
17028                Err(e) => eprintln!(
17029                    "refit: FAILED to write {path} ({} MB): {e}",
17030                    bytes.len() / 1_000_000
17031                ),
17032            }
17033        };
17034        w("gss", &acc.gss);
17035        w("ya", &acc.ya);
17036        println!(
17037            "refit L{li}: {} support, {} tokens, hidden {}",
17038            acc.support.len(),
17039            acc.tokens,
17040            acc.hidden
17041        );
17042        n += 1;
17043    }
17044    n
17045}
17046
17047/// `CMF_FFN_ADUMP=<prefix>` — append every probed token's FFN activation
17048/// row to `<prefix>.<layer>.f16`. The co-activation record: which
17049/// neurons fire together, which is what a tube has to group if a token
17050/// is ever going to open one tube instead of sixteen.
17051fn adump_row(li: usize, g: &[f32]) {
17052    use std::io::Write as _;
17053    static FILES: std::sync::OnceLock<
17054        Option<(
17055            String,
17056            std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
17057        )>,
17058    > = std::sync::OnceLock::new();
17059    let Some((prefix, map)) = FILES
17060        .get_or_init(|| {
17061            std::env::var("CMF_FFN_ADUMP")
17062                .ok()
17063                .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
17064        })
17065        .as_ref()
17066    else {
17067        return;
17068    };
17069    // `CMF_FFN_ADUMP_FROM/_TO` narrow the dump to a layer span, so a big
17070    // calibration run fits on disk in a few passes instead of one.
17071    static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
17072    let (from, to) = *SPAN.get_or_init(|| {
17073        let g = |k: &str, d: usize| {
17074            std::env::var(k)
17075                .ok()
17076                .and_then(|v| v.parse().ok())
17077                .unwrap_or(d)
17078        };
17079        (
17080            g("CMF_FFN_ADUMP_FROM", 0),
17081            g("CMF_FFN_ADUMP_TO", usize::MAX),
17082        )
17083    });
17084    if li < from || li > to {
17085        return;
17086    }
17087    let mut map = map.lock().unwrap();
17088    let f = map.entry(li).or_insert_with(|| {
17089        std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
17090    });
17091    let mut bytes = Vec::with_capacity(g.len() * 2);
17092    for v in g {
17093        bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
17094    }
17095    let _ = f.write_all(&bytes);
17096}
17097
17098/// `CMF_FFN_ORACLE_TOPK` — keep only the k largest |silu(g)·u| of each
17099/// token and zero the rest. Not a serving mode: it is the CEILING of
17100/// contextual sparsity — what a per-token router would be chasing —
17101/// measured by cheating, since the selection reads the very activations
17102/// it would have to predict.
17103fn oracle_topk() -> usize {
17104    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17105    *K.get_or_init(|| {
17106        std::env::var("CMF_FFN_ORACLE_TOPK")
17107            .ok()
17108            .and_then(|v| v.parse().ok())
17109            .unwrap_or(0)
17110    })
17111}
17112
17113/// `CMF_FFN_GATE_TOPK` — the REALIZABLE cousin of the oracle: rank the
17114/// neurons by their gate alone (which the kernel has computed anyway
17115/// before it reads `up`), keep the k best, and drop the rest. Every
17116/// dropped neuron's `up` row and `down` column stay unread, so this is
17117/// the sparsity a serving path can actually take without a router.
17118fn gate_topk() -> usize {
17119    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17120    *K.get_or_init(|| {
17121        std::env::var("CMF_FFN_GATE_TOPK")
17122            .ok()
17123            .and_then(|v| v.parse().ok())
17124            .unwrap_or(0)
17125    })
17126}
17127
17128/// `CMF_FFN_GATE_BLOCK` — select in blocks of B neurons instead of one
17129/// by one. A scattered per-neuron choice cannot be read efficiently (a
17130/// row at a time, no prefetch runway); a block of 32 is a contiguous
17131/// 32-row slab of `up` and of the transposed `down`, which the ordinary
17132/// kernels stream. The question the measurement answers is what the
17133/// block costs in quality.
17134fn gate_block() -> usize {
17135    static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17136    *B.get_or_init(|| {
17137        std::env::var("CMF_FFN_GATE_BLOCK")
17138            .ok()
17139            .and_then(|v| v.parse().ok())
17140            .unwrap_or(1)
17141    })
17142}
17143
17144/// Zero all but the `k` largest BLOCKS (by summed square) of a row.
17145fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
17146    let n = g.len();
17147    let nb = n.div_ceil(block);
17148    let kb = (keep_n.div_ceil(block)).clamp(1, nb);
17149    if kb >= nb {
17150        return;
17151    }
17152    let mut score: Vec<f32> = (0..nb)
17153        .map(|b| {
17154            g[b * block..((b + 1) * block).min(n)]
17155                .iter()
17156                .map(|v| v * v)
17157                .sum::<f32>()
17158        })
17159        .collect();
17160    let mut ord = score.clone();
17161    let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
17162        b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17163    });
17164    let thr = *kth;
17165    for b in 0..nb {
17166        if score[b] < thr {
17167            g[b * block..((b + 1) * block).min(n)].fill(0.0);
17168        }
17169    }
17170    score.clear();
17171}
17172
17173/// Zero all but the `k` largest magnitudes of one token's activation row.
17174fn keep_top_k(g: &mut [f32], k: usize) {
17175    if gate_block() > 1 {
17176        return keep_top_blocks(g, k, gate_block());
17177    }
17178    let n = g.len();
17179    if k == 0 || k >= n {
17180        return;
17181    }
17182    let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
17183    let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
17184        b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17185    });
17186    let thr = *kth;
17187    for v in g.iter_mut() {
17188        if v.abs() < thr {
17189            *v = 0.0;
17190        }
17191    }
17192}
17193
17194/// `CMF_FFN_PROBE_SQ` — accumulate Σa², so the dump divided by the token
17195/// count and square-rooted is the RMS activation trace Patent 12 weights
17196/// its matrices by.
17197fn probe_sq() -> bool {
17198    static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17199    *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
17200}
17201
17202/// `CMF_FFN_PROBE_SIGNED` — accumulate the SIGNED activation sum
17203/// instead of its magnitude: what a dropped neuron contributes ON
17204/// AVERAGE, which is the bias a narrowed FFN can add back for free.
17205fn probe_signed() -> bool {
17206    static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17207    *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
17208}
17209
17210/// `CMF_FFN_MEANFILL=<file>` — a masked-out neuron contributes its MEAN
17211/// activation instead of zero (`u32 layers, u32 inter, f32[…]`, the mass
17212/// dump layout, holding per-neuron means). Dropping a neuron outright
17213/// also drops its average contribution, which shifts the layer output by
17214/// a constant; filling the mean back is one add per layer and costs no
17215/// bytes off the bus. This is the measurement arm — in a tube file the
17216/// same correction ships as a per-task bias vector.
17217fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
17218    static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
17219    M.get_or_init(|| {
17220        let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
17221        let b = std::fs::read(&p).ok()?;
17222        let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
17223        let vals: Vec<f32> = b[8..]
17224            .chunks_exact(4)
17225            .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
17226            .collect();
17227        eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
17228        Some((inter, vals))
17229    })
17230    .as_ref()
17231}
17232
17233/// `CMF_FFN_PROBE_TOPK` — 0 (default) = accumulate mass, k>0 = count
17234/// how often a neuron lands in a token's top k.
17235fn probe_topk() -> usize {
17236    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17237    *K.get_or_init(|| {
17238        std::env::var("CMF_FFN_PROBE_TOPK")
17239            .ok()
17240            .and_then(|v| v.parse().ok())
17241            .unwrap_or(0)
17242    })
17243}
17244
17245thread_local! {
17246    /// DTG-MA activation probe: per-layer per-neuron Σ|silu(g)·u|
17247    /// accumulator, alive only during `Pipeline::probe_ffn_mass`.
17248    static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
17249        const { std::cell::RefCell::new(None) };
17250}
17251
17252/// Per-token structured sparsity, paid for in bytes.
17253///
17254/// The gate is the cheapest third of an FFN and it already says which
17255/// neurons matter: `silu(gate)` near zero means the neuron contributes
17256/// nothing whatever `up` says. So compute every gate, keep the `k`
17257/// loudest, and read ONLY those neurons' `up` rows and `down` rows —
17258/// the latter needs `down_proj` stored transposed, otherwise a neuron's
17259/// down weights are a strided column and "reading only those" costs a
17260/// full cache line each.
17261///
17262/// Returns `None` when the file has no transposed `down` (the caller
17263/// then runs the ordinary dense path).
17264fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
17265    // The scatter path reads individual rows/columns and cannot express the
17266    // per-matrix signed FWHT boundary.  Let the descriptor-aware dense path
17267    // handle Prism files rather than silently running an unrotated sparse
17268    // approximation.
17269    if d.gate_proj.has_prism_contract()
17270        || d.up_proj.has_prism_contract()
17271        || d.down_proj.has_prism_contract()
17272    {
17273        return None;
17274    }
17275    let dt = d.down_t.as_ref()?;
17276    let inter = d.gate_proj.rows();
17277    let hidden = dt.cols();
17278    if k == 0 || k >= inter || d.act != Act::Silu {
17279        return None;
17280    }
17281    DYN_SCRATCH.with(|sc| {
17282        let mut sc = sc.borrow_mut();
17283        let DynScratch {
17284            g,
17285            mag,
17286            live,
17287            parts,
17288        } = &mut *sc;
17289        g.resize(inter, 0.0);
17290        d.gate_proj.matvec(x, g, pool);
17291        for v in g.iter_mut() {
17292            *v = inference::silu(*v);
17293        }
17294        // The k-th largest |silu(gate)| is the threshold; ties keep more,
17295        // which is the safe side.
17296        mag.clear();
17297        mag.extend(g.iter().map(|v| v.abs()));
17298        let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
17299            b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17300        });
17301        let thr = *kth;
17302        live.clear();
17303        live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
17304        let mut out = vec![0.0f32; hidden];
17305        match pool {
17306            Some(p) if live.len() >= 64 => {
17307                let nw = p.n_workers() + 1;
17308                parts.clear();
17309                parts.resize(nw * hidden, 0.0);
17310                let ptr = SendMut(parts.as_mut_ptr());
17311                let n = live.len();
17312                let live_ref: &[u32] = live;
17313                let g_ref: &[f32] = g;
17314                p.run(&|w, workers| {
17315                    let chunk = n.div_ceil(workers);
17316                    let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
17317                    if s >= e {
17318                        return;
17319                    }
17320                    WORKER_SCRATCH.with(|ws| {
17321                        let mut ws = ws.borrow_mut();
17322                        let [scratch, acc] = &mut *ws;
17323                        scratch.resize(hidden.max(x.len()), 0.0);
17324                        acc.clear();
17325                        acc.resize(hidden, 0.0);
17326                        for (o, &nrm) in live_ref[s..e].iter().enumerate() {
17327                            // One neuron of runway: the next row's lines
17328                            // start moving while this one is multiplied.
17329                            if let Some(&nx) = live_ref[s..e].get(o + 1) {
17330                                d.up_proj.prefetch_row(nx as usize);
17331                                dt.prefetch_row(nx as usize);
17332                            }
17333                            let idx = nrm as usize;
17334                            let up = d.up_proj.row_dot(idx, x, scratch);
17335                            let a = g_ref[idx] * up;
17336                            if a != 0.0 {
17337                                dt.add_row_scaled(idx, a, acc, scratch);
17338                            }
17339                        }
17340                        for (j, v) in acc.iter().enumerate() {
17341                            unsafe { *ptr.at(w * hidden + j) = *v };
17342                        }
17343                    });
17344                });
17345                for w in 0..nw {
17346                    for (j, o) in out.iter_mut().enumerate() {
17347                        *o += parts[w * hidden + j];
17348                    }
17349                }
17350            }
17351            _ => {
17352                WORKER_SCRATCH.with(|ws| {
17353                    let mut ws = ws.borrow_mut();
17354                    let [scratch, _acc] = &mut *ws;
17355                    scratch.resize(hidden.max(x.len()), 0.0);
17356                    for &nrm in live.iter() {
17357                        let idx = nrm as usize;
17358                        let up = d.up_proj.row_dot(idx, x, scratch);
17359                        let a = g[idx] * up;
17360                        if a != 0.0 {
17361                            dt.add_row_scaled(idx, a, &mut out, scratch);
17362                        }
17363                    }
17364                });
17365            }
17366        }
17367        Some(out)
17368    })
17369}
17370
17371/// Caller-side scratch of the dynamic path — one allocation per thread,
17372/// not one per layer per token (that alone cost a third of the decode).
17373struct DynScratch {
17374    g: Vec<f32>,
17375    mag: Vec<f32>,
17376    live: Vec<u32>,
17377    parts: Vec<f32>,
17378}
17379
17380thread_local! {
17381    static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
17382        std::cell::RefCell::new(DynScratch {
17383            g: Vec::new(),
17384            mag: Vec::new(),
17385            live: Vec::new(),
17386            parts: Vec::new(),
17387        })
17388    };
17389    /// Pool-worker scratch: the row buffer and this worker's partial sum.
17390    static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
17391        const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
17392}
17393
17394/// `dense_ffn_cpu` with a per-visit mask landing on the activations —
17395/// the masked-inference fast path's decode arm. Full fused quant
17396/// compute, closed neurons zeroed before down: arithmetically the
17397/// pruned network, no dequant, no weight bytes touched.
17398fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
17399    let inter = d.gate_proj.rows();
17400    FFN_SCRATCH.with(|s| {
17401        let mut s = s.borrow_mut();
17402        let [g, u, ..] = &mut *s;
17403        g.resize(inter, 0.0);
17404        if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
17405            // g holds silu(gate)·up.
17406        } else {
17407            u.resize(inter, 0.0);
17408            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
17409            for i in 0..inter {
17410                g[i] = d.act.combine(g[i], u[i]);
17411            }
17412        }
17413        zero_masked_cols(g, 1, inter, mask_row);
17414        let mut out = attention::take_buf(d.down_proj.rows());
17415        d.down_proj.matvec(g, &mut out, pool);
17416        out
17417    })
17418}
17419
17420/// Dense FFN as one GPU submission via the MoE block path (single
17421/// expert, weight 1.0): gate → silu·up → down chained in one command
17422/// buffer, intermediate activations device-resident. None → weights
17423/// not q8-mapped in the primary shard / over the VRAM budget / backend
17424/// refusal → honest CPU path.
17425fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
17426    if d.gate_proj.has_prism_contract()
17427        || d.up_proj.has_prism_contract()
17428        || d.down_proj.has_prism_contract()
17429    {
17430        return None;
17431    }
17432    // The GPU block hardcodes SiLU; GeLU FFNs (Gemma) stay on CPU.
17433    if d.act != Act::Silu {
17434        return None;
17435    }
17436    // Threshold: tiny FFNs are not worth a submission (q1 excepted —
17437    // see the caller's gate).
17438    if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
17439        return None;
17440    }
17441    let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
17442    let mut model_ref = None;
17443    moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
17444    let model = model_ref?;
17445    let hidden = jobs[0].down.1;
17446    let mut out = attention::take_buf(hidden);
17447    if crate::gpu::moe_block(&model, &jobs, &mut out) {
17448        Some(out)
17449    } else {
17450        let mut out = out;
17451        attention::recycle_buf(&mut out);
17452        None
17453    }
17454}
17455
17456/// q8-mapped primary-shard tensor parts for a GPU job: q8_2f carries
17457/// its column field, q8_row runs with empty col slices (the backend
17458/// skips the multiply). Shared by the MoE block and the dense-FFN
17459/// single-job path.
17460#[allow(clippy::type_complexity)]
17461#[allow(clippy::type_complexity)]
17462pub(crate) fn moe_parts(
17463    t: &QTensor,
17464) -> Option<(
17465    &std::sync::Arc<cortiq_core::CmfModel>,
17466    usize,
17467    usize,
17468    usize,
17469    &[f32],
17470    &[f32],
17471    bool,
17472    bool,
17473    bool,
17474)> {
17475    match t {
17476        QTensor::Mapped {
17477            model,
17478            idx,
17479            dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
17480            rows,
17481            cols,
17482            row_scale,
17483            col_field,
17484            ..
17485        } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
17486            model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
17487        )),
17488        // q1: tile-embedded scales — empty rs/col slices, raw xs.
17489        QTensor::Mapped {
17490            model,
17491            idx,
17492            dtype: cortiq_core::TensorDtype::Q1,
17493            rows,
17494            cols,
17495            ..
17496        } => Some((
17497            model,
17498            *idx,
17499            *rows,
17500            *cols,
17501            &[][..],
17502            &[][..],
17503            true,
17504            false,
17505            false,
17506        )),
17507        // q4_tiled: 18-byte tiles with embedded f16 scales — raw xs.
17508        QTensor::Mapped {
17509            model,
17510            idx,
17511            dtype: cortiq_core::TensorDtype::Q4Tiled,
17512            rows,
17513            cols,
17514            ..
17515        } => Some((
17516            model,
17517            *idx,
17518            *rows,
17519            *cols,
17520            &[][..],
17521            &[][..],
17522            false,
17523            true,
17524            false,
17525        )),
17526        // q4tp: same raw-xs contract, different stride and scale plane.
17527        QTensor::Mapped {
17528            model,
17529            idx,
17530            dtype: cortiq_core::TensorDtype::Q4TiledP,
17531            rows,
17532            cols,
17533            ..
17534        } => Some((
17535            model,
17536            *idx,
17537            *rows,
17538            *cols,
17539            &[][..],
17540            &[][..],
17541            false,
17542            true,
17543            false,
17544        )),
17545        // q2tp: the 2-bit expert plane of the mixed profile — q4 family
17546        // for stride bookkeeping, flagged q2 so the trio validation can
17547        // demand a q4tp down.
17548        QTensor::Mapped {
17549            model,
17550            idx,
17551            dtype: cortiq_core::TensorDtype::Q2TiledP,
17552            rows,
17553            cols,
17554            ..
17555        } => Some((
17556            model,
17557            *idx,
17558            *rows,
17559            *cols,
17560            &[][..],
17561            &[][..],
17562            false,
17563            true,
17564            true,
17565        )),
17566        _ => None,
17567    }
17568}
17569
17570/// Map a MoE onto the Metal token graph's contract: f32 router, a
17571/// shared expert (gated — Qwen — or ungated at weight 1 — DeepSeek-V3 /
17572/// HunYuan hy_v3), softmax or sigmoid scores with an optional selection
17573/// bias and routed scale, experts uniformly q4tp (or the mixed profile:
17574/// q2tp gate/up over a q4tp down). τ routers, masks, per-expert scales
17575/// and Gemma's router-input norm refuse here — those semantics stay on
17576/// the CPU path.
17577#[cfg(target_os = "macos")]
17578fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
17579    if m.router_input_norm
17580        || m.route_tau.is_some()
17581        || m.mask.is_some()
17582        || m.per_expert_scale.is_some()
17583        || m.experts.is_empty()
17584        || m.top_k == 0
17585        || m.resonance.is_some()
17586    {
17587        return None;
17588    }
17589    // The select kernel always fills the shared slot: a model without a
17590    // shared expert (LFM2-MoE) stays on the CPU path here.
17591    let (sh, sg) = match &m.shared {
17592        Some((sh, sg)) => (sh, sg.as_ref()),
17593        None => return None,
17594    };
17595    let (rf, rr, rc) = m.router.f32_parts()?;
17596    if rr != m.experts.len() || rc != hidden {
17597        return None;
17598    }
17599    let shared_gated = sg.is_some();
17600    let sf = match sg {
17601        Some(sg) => {
17602            let (sf, sr, sc) = sg.f32_parts()?;
17603            if sr * sc != hidden {
17604                return None;
17605            }
17606            sf
17607        }
17608        // Ungated: the router's first row stands in for the gate matvec
17609        // (its logit is never read — the kernel pins weight 1).
17610        None => &rf[..hidden],
17611    };
17612    if let Some(b) = &m.expert_bias {
17613        if b.len() != m.experts.len() {
17614            return None;
17615        }
17616    }
17617    let inter = m.experts[0].gate_proj.rows();
17618    // The first expert's gate decides the profile; every trio (shared
17619    // included) must agree — the jobs ladder flips ONE kernel for all.
17620    let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
17621    let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
17622        if e.act != Act::Silu
17623            || e.gate_proj.rows() != inter
17624            || e.gate_proj.cols() != hidden
17625            || e.up_proj.rows() != inter
17626            || e.up_proj.cols() != hidden
17627            || e.down_proj.rows() != hidden
17628            || e.down_proj.cols() != inter
17629        {
17630            return None;
17631        }
17632        let pick = |t: &QTensor| -> Option<usize> {
17633            if gu_q2 {
17634                t.mapped_q2tp().map(|(_, i)| i)
17635            } else {
17636                t.mapped_q4tp().map(|(_, i)| i)
17637            }
17638        };
17639        Some((
17640            pick(&e.gate_proj)?,
17641            pick(&e.up_proj)?,
17642            e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
17643        ))
17644    };
17645    let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
17646    let shared = trio(sh)?;
17647    Some(crate::gpu::GpuMoe {
17648        router: rf,
17649        sgate: sf,
17650        experts,
17651        shared,
17652        n_exp: m.experts.len(),
17653        top_k: m.top_k,
17654        inter,
17655        norm_topk: m.norm_topk_prob,
17656        route_scale: m.routed_scaling,
17657        gu_q2,
17658        sigmoid: m.router_sigmoid,
17659        bias: m.expert_bias.as_deref(),
17660        shared_gated,
17661    })
17662}
17663
17664/// Build one gate/up/down GPU job from three tensors. `moe_push_job` is the
17665/// DenseFfn-shaped caller; architectures that keep their experts in their own
17666/// structs (DeepSeek-V4) come here directly.
17667pub(crate) fn moe_push_job_parts<'a>(
17668    gate: &'a QTensor,
17669    up: &'a QTensor,
17670    down: &'a QTensor,
17671    x: &[f32],
17672    w: f32,
17673    swiglu_limit: f32,
17674    jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17675    model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17676) -> Option<()> {
17677    use crate::qtensor::prescale;
17678    let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
17679    let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
17680    let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
17681    if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17682        return None; // mixed-dtype trio — honest CPU path
17683    }
17684    // The 2-bit profile is gate/up q2tp over a PLAIN q4tp down; any other
17685    // 2-bit arrangement stays on the CPU.
17686    if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
17687        return None;
17688    }
17689    if !gq2 && dq2 {
17690        return None;
17691    }
17692    model_ref.get_or_insert_with(|| gm.clone());
17693    let dt = |cf: &[f32]| {
17694        if cf.is_empty() {
17695            cortiq_core::TensorDtype::Q8Row
17696        } else {
17697            cortiq_core::TensorDtype::Q8_2f
17698        }
17699    };
17700    jobs.push(crate::gpu::MoeJob {
17701        gate: (gi, gr, gc, grs),
17702        up: (ui, ur, uc, urs),
17703        down: (di, dr, dc, drs),
17704        xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
17705        xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
17706        down_col: dcf,
17707        w,
17708        q1: gq1,
17709        q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
17710        q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
17711        gu_q2: gq2,
17712        swiglu_limit,
17713    });
17714    Some(())
17715}
17716
17717/// Build one gate/up/down GPU job (see `moe_parts`).
17718fn moe_push_job<'a>(
17719    d: &'a DenseFfn,
17720    x: &[f32],
17721    w: f32,
17722    jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17723    model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17724) -> Option<()> {
17725    use crate::qtensor::prescale;
17726    if d.act != Act::Silu {
17727        return None; // GPU block hardcodes SiLU
17728    }
17729    let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
17730    let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
17731    let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
17732    if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17733        return None; // mixed-dtype trio — honest CPU path
17734    }
17735    if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
17736        return None;
17737    }
17738    if !gq2 && dq2 {
17739        return None;
17740    }
17741    model_ref.get_or_insert_with(|| gm.clone());
17742    let gdt = if gcf.is_empty() {
17743        cortiq_core::TensorDtype::Q8Row
17744    } else {
17745        cortiq_core::TensorDtype::Q8_2f
17746    };
17747    let udt = if ucf.is_empty() {
17748        cortiq_core::TensorDtype::Q8Row
17749    } else {
17750        cortiq_core::TensorDtype::Q8_2f
17751    };
17752    jobs.push(crate::gpu::MoeJob {
17753        gate: (gi, gr, gc, grs),
17754        up: (ui, ur, uc, urs),
17755        down: (di, dr, dc, drs),
17756        xs_gate: prescale(x, gcf, gdt).into_owned(),
17757        xs_up: prescale(x, ucf, udt).into_owned(),
17758        down_col: dcf,
17759        w,
17760        q1: gq1,
17761        q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
17762        q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
17763        gu_q2: gq2,
17764        swiglu_limit: 0.0,
17765    });
17766    Some(())
17767}
17768
17769/// Sparse dense-FFN directly on QUANTIZED weights (mask × mmap): reads
17770/// ONLY the active neurons' gate/up rows and down columns from the mmap
17771/// — no full-matrix dequant, no f32 model copy. This is what lets a
17772/// masked big model run at quantized RSS (the historical mask path
17773/// forced the whole model to f32). Semantics identical to the f32
17774/// sparse path within quant tolerance.
17775fn sparse_ffn_quant(
17776    d: &DenseFfn,
17777    x: &[f32],
17778    active: &[u16],
17779    hidden: usize,
17780    pool: Option<&Pool>,
17781) -> Vec<f32> {
17782    let n = active.len();
17783    let inter = d.gate_proj.rows();
17784    let mut act = vec![0.0f32; n];
17785    // Scratch is needed if EITHER projection is group-packed (q4/vbit);
17786    // gate/up normally share a dtype but sizing on both is robust.
17787    let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
17788    let compute = |ai: usize| -> f32 {
17789        let idx = active[ai] as usize;
17790        if idx >= inter {
17791            return 0.0; // defensive parity with the f32 sparse path
17792        }
17793        let mut s = if need_scratch {
17794            vec![0.0f32; hidden]
17795        } else {
17796            Vec::new()
17797        };
17798        let gate = d.gate_proj.row_dot(idx, x, &mut s);
17799        let up = d.up_proj.row_dot(idx, x, &mut s);
17800        d.act.combine(gate, up)
17801    };
17802    match pool {
17803        Some(p) if n >= 256 => {
17804            let ptr = SendMut(act.as_mut_ptr());
17805            p.run(&|widx, nw| {
17806                let chunk = n.div_ceil(nw);
17807                let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
17808                for ai in s..e {
17809                    unsafe { *ptr.at(ai) = compute(ai) };
17810                }
17811            });
17812        }
17813        _ => {
17814            for (ai, a) in act.iter_mut().enumerate() {
17815                *a = compute(ai);
17816            }
17817        }
17818    }
17819    // Scatter through active down columns (reads only those columns).
17820    let mut out = vec![0.0f32; hidden];
17821    for (ai, &idx) in active.iter().enumerate() {
17822        let w = act[ai];
17823        if w.abs() >= 1e-12 && (idx as usize) < inter {
17824            d.down_proj.add_col_scaled(idx as usize, w, &mut out);
17825        }
17826    }
17827    out
17828}
17829
17830/// Test-only re-export of the private sparse-quant FFN (mask × mmap gate).
17831#[doc(hidden)]
17832pub fn sparse_ffn_quant_for_test(
17833    d: &DenseFfn,
17834    x: &[f32],
17835    active: &[u16],
17836    hidden: usize,
17837) -> Vec<f32> {
17838    sparse_ffn_quant(d, x, active, hidden, None)
17839}
17840
17841/// Dequantize a DenseFfn's three matrices to f32 (transient; only the
17842/// q4/vbit-masked fallback uses it — the memory-lean path is
17843/// sparse_ffn_quant). Reuses row_f32 row-by-row.
17844fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
17845    let deq = |t: &QTensor| -> Vec<f32> {
17846        let (rows, cols) = (t.rows(), t.cols());
17847        let mut out = vec![0.0f32; rows * cols];
17848        for r in 0..rows {
17849            t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
17850        }
17851        out
17852    };
17853    (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
17854}
17855
17856/// Pointer wrapper for the worker-pool scatter (same pattern as qtensor).
17857struct SendMut(*mut f32);
17858unsafe impl Send for SendMut {}
17859unsafe impl Sync for SendMut {}
17860impl SendMut {
17861    #[inline]
17862    // Deliberate unsynchronized scatter: pool workers write disjoint indices
17863    // in parallel, so returning `&mut` from `&self` is intentional here.
17864    #[allow(clippy::mut_from_ref)]
17865    unsafe fn at(&self, i: usize) -> &mut f32 {
17866        unsafe { &mut *self.0.add(i) }
17867    }
17868}
17869
17870/// Router → (selected experts in torch.topk order, per-expert score
17871/// vector, normalizer). The final weight of expert `e` is `p[e] / wsum`.
17872///
17873/// Two regimes share this. Qwen: softmax over ALL experts, top-k of the
17874/// probabilities, optional renorm — `router_sigmoid=false`, no bias,
17875/// scale 1 → bit-identical to the historical path. LFM2-MoE /
17876/// DeepSeek-V3 `noaux_tc`: per-expert sigmoid scores, an optional
17877/// selection bias (top-k CHOICE only; weights stay unbiased), a 1e-6 renorm
17878/// floor and a routed scale. Architectures whose reference uses a different
17879/// sigmoid denominator floor (for example GLM-5's `1e-20`) call
17880/// [`moe_route_with_eps`] directly; the historical generic path remains
17881/// unchanged.
17882pub(crate) fn moe_route(
17883    logits: &[f32],
17884    m: &MoeFfn,
17885    allowed: Option<&[bool]>,
17886) -> (Vec<usize>, Vec<f32>, f32) {
17887    moe_route_with_eps(logits, m, allowed, 1e-6)
17888}
17889
17890/// Router implementation with an explicit sigmoid renormalization floor.
17891///
17892/// GLM-5.3's source computes `sum(selected_scores) + 1e-20`; using the
17893/// generic 1e-6 floor there is not a harmless tolerance difference when all
17894/// logits are very negative: it collapses the routed branch toward zero
17895/// instead of normalizing the selected experts. Keeping the epsilon parameter
17896/// here avoids changing the established Qwen/LFM2 contract while allowing
17897/// each architecture to preserve its own numerical semantics.
17898pub(crate) fn moe_route_with_eps(
17899    logits: &[f32],
17900    m: &MoeFfn,
17901    allowed: Option<&[bool]>,
17902    sigmoid_denom_eps: f32,
17903) -> (Vec<usize>, Vec<f32>, f32) {
17904    let ne = logits.len();
17905    // Expert restriction: the static env mask (CMF_MOE_MASK) AND the
17906    // active task mask's expert fields (spec §5) both narrow the
17907    // candidate set; selection happens over the admitted experts only.
17908    // With norm_topk the kept weights renormalize below; without it
17909    // the excluded mass is honestly dropped.
17910    let admit = |e: usize| {
17911        m.mask.as_ref().is_none_or(|mk| mk[e])
17912            && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
17913    };
17914    // The resonance router (spec §9.5.1) selects by the RAW score: the
17915    // trainer (`resonance_winner`) and the resident graph
17916    // (`embryo_core_route_pick`) take the first maximum of the scores
17917    // and run the winner with weight 1.0. Selecting through the softmax
17918    // instead is not the same decision: `exp(l − max)` rounds two scores
17919    // closer than 2^-25 (possible below |score| 0.25) to the same 1.0, and
17920    // the lower index would take a token whose score is strictly smaller
17921    // — the trainer's trace and the graph would disagree with this path.
17922    // `−∞` (outside the shell) never wins; with no finite admitted expert
17923    // the generic path below degrades to uniform. The winner's
17924    // probability is 1.0 by construction (a one-hot `p`), so its
17925    // renormalized weight is `routed_scaling` on both norm_topk settings.
17926    if m.resonance.is_some() && m.top_k == 1 {
17927        let mut best: Option<usize> = None;
17928        for e in (0..ne).filter(|&e| admit(e)) {
17929            let l = logits[e];
17930            if l == f32::NEG_INFINITY || l.is_nan() {
17931                continue;
17932            }
17933            if best.is_none_or(|b| l > logits[b]) {
17934                best = Some(e);
17935            }
17936        }
17937        if let Some(b) = best {
17938            let mut p = vec![0.0f32; ne];
17939            p[b] = 1.0;
17940            return (vec![b], p, 1.0 / m.routed_scaling);
17941        }
17942    }
17943    // A `−∞` logit (a grown expert outside its shell, `Resonance::scores`)
17944    // takes probability 0 on both paths: sigmoid(−∞) = 0, exp(−∞ − max) = 0
17945    // — top-1 is the best FINITE expert, its renormalized weight exactly
17946    // 1.0. Every expert at −∞ cannot happen (trunk experts have no shell);
17947    // should it, the softmax would be NaN, so it degrades to uniform.
17948    let p: Vec<f32> = if m.router_sigmoid {
17949        logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
17950    } else {
17951        let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
17952        if mx == f32::NEG_INFINITY {
17953            vec![1.0 / ne.max(1) as f32; ne]
17954        } else {
17955            let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
17956            let s: f32 = e.iter().sum();
17957            for v in &mut e {
17958                *v /= s;
17959            }
17960            e
17961        }
17962    };
17963    let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
17964    // Descending by selection score, lower index wins ties (torch.topk).
17965    match &m.expert_bias {
17966        Some(b) => idx.sort_unstable_by(|&x, &y| {
17967            (p[y] + b[y])
17968                .partial_cmp(&(p[x] + b[x]))
17969                .unwrap()
17970                .then(x.cmp(&y))
17971        }),
17972        None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
17973    }
17974    idx.truncate(m.top_k);
17975    // Adaptive τ-routing: trim the tail experts once the kept mass is
17976    // enough. wsum below renormalizes over the KEPT set, so the output
17977    // stays a proper weighted average.
17978    if let Some(tau) = m.route_tau {
17979        let total: f32 = idx.iter().map(|&e| p[e]).sum();
17980        if total > 0.0 {
17981            let mut acc = 0.0f32;
17982            let mut keep = idx.len();
17983            for (i, &e) in idx.iter().enumerate() {
17984                acc += p[e];
17985                if acc >= tau * total {
17986                    keep = i + 1;
17987                    break;
17988                }
17989            }
17990            idx.truncate(keep);
17991        }
17992    }
17993    let wsum: f32 = if m.norm_topk_prob {
17994        let s: f32 = idx.iter().map(|&e| p[e]).sum();
17995        // Sigmoid routers use their architecture's reference floor; the
17996        // softmax path's probs already sum near 1, so it stays exactly as
17997        // before.
17998        (if m.router_sigmoid {
17999            s + sigmoid_denom_eps
18000        } else {
18001            s
18002        }) / m.routed_scaling
18003    } else {
18004        1.0 / m.routed_scaling
18005    };
18006    (idx, p, wsum)
18007}
18008
18009/// See the call site: one `layer:e1,e2,…` line per routed token.
18010fn moe_trace(idx: &[usize]) {
18011    moe_trace_at(crate::gpu::cur_layer() as i32, idx)
18012}
18013
18014/// The same, for callers that know their layer (DSV4 owns its layers and
18015/// never sets the pipeline's current-layer marker).
18016pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
18017    use std::io::Write;
18018    static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
18019        std::sync::OnceLock::new();
18020    let Some(f) = F.get_or_init(|| {
18021        let p = std::env::var("CMF_MOE_TRACE").ok()?;
18022        Some(std::sync::Mutex::new(
18023            std::fs::OpenOptions::new()
18024                .create(true)
18025                .append(true)
18026                .open(p)
18027                .ok()?,
18028        ))
18029    }) else {
18030        return;
18031    };
18032    let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
18033    let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
18034}
18035
18036/// MoE FFN: router → top-k experts (see `moe_route`). Only selected
18037/// experts' pages are touched in mmap.
18038pub(crate) fn moe_ffn(
18039    m: &MoeFfn,
18040    x: &[f32],
18041    pool: Option<&Pool>,
18042    allowed: Option<&[bool]>,
18043) -> Vec<f32> {
18044    let r = moe_ffn_route(m, x, pool, allowed);
18045    moe_ffn_experts(m, x, &r, pool)
18046}
18047
18048/// One token's host route through a MoE layer: the chosen experts in
18049/// selection order, the per-expert scores and the normalizer (see
18050/// `moe_route`), plus the raw router logits.
18051pub(crate) struct MoeRoute {
18052    pub idx: Vec<usize>,
18053    pub p: Vec<f32>,
18054    pub wsum: f32,
18055    pub logits: Vec<f32>,
18056}
18057
18058/// The routing half of `moe_ffn`, shared by every executor of the chosen
18059/// experts (the host/per-op path below and the MiMo dynamic device cache,
18060/// `crate::mimo_moe`): activation accounting, router logits, `moe_route`,
18061/// the selection statistics and the `CMF_MOE_TRACE` line — so switching
18062/// executors can never change which experts a token gets.
18063pub(crate) fn moe_ffn_route(
18064    m: &MoeFfn,
18065    x: &[f32],
18066    pool: Option<&Pool>,
18067    allowed: Option<&[bool]>,
18068) -> MoeRoute {
18069    accumulate_act(m, x, 1);
18070    let ne = m.experts.len();
18071    let mut logits = vec![0.0f32; ne];
18072    match &m.resonance {
18073        Some(r) => r.scores(x, &mut logits),
18074        None => m.router.matvec(x, &mut logits, pool),
18075    }
18076    let (idx, p, wsum) = moe_route(&logits, m, allowed);
18077    {
18078        let mut st = m.stats.borrow_mut();
18079        if st.len() < ne {
18080            st.resize(ne, 0);
18081        }
18082        for &e in &idx {
18083            st[e] += 1;
18084        }
18085    }
18086    // `CMF_MOE_TRACE=<file>`: append one line per (layer, token) with the
18087    // selected expert ids. The cumulative `stats` above answer "which
18088    // experts are popular"; a residency design needs the question they
18089    // cannot answer — whether CONSECUTIVE tokens reuse experts (the
18090    // temporal locality an LRU cache lives on, FreeToken §4).
18091    moe_trace(&idx);
18092    MoeRoute {
18093        idx,
18094        p,
18095        wsum,
18096        logits,
18097    }
18098}
18099
18100/// The expert half of `moe_ffn`: run a route's experts on the per-op GPU
18101/// block or the host.
18102pub(crate) fn moe_ffn_experts(
18103    m: &MoeFfn,
18104    x: &[f32],
18105    r: &MoeRoute,
18106    pool: Option<&Pool>,
18107) -> Vec<f32> {
18108    let (idx, p, wsum) = (&r.idx, &r.p, r.wsum);
18109    // D5: the whole layer MoE block in one GPU command buffer (experts — the
18110    // same mmap via a no-copy buffer; intermediate activations on the GPU).
18111    // Same Ffn probe class as the dense chain: one submit per layer
18112    // either wins on this driver stack or it doesn't.
18113    if crate::gpu::enabled_here() {
18114        match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
18115            crate::gpu::ProbeArm::Gpu => {
18116                let t0 = std::time::Instant::now();
18117                if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
18118                    crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
18119                    return out;
18120                }
18121            }
18122            crate::gpu::ProbeArm::CpuTimed => {
18123                let t0 = std::time::Instant::now();
18124                let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
18125                crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
18126                return out;
18127            }
18128            crate::gpu::ProbeArm::Cpu => {
18129                return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
18130            }
18131        }
18132    }
18133    moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
18134}
18135
18136/// One MoE token through the MiMo expert bank (`crate::mimo_moe`), or —
18137/// when the bank does not serve it — through the host path with the SAME
18138/// route, so the routing statistics and `CMF_MOE_TRACE` see it once.
18139fn moe_ffn_banked(
18140    slot: &mut crate::mimo_moe::Slot,
18141    li: usize,
18142    m: &MoeFfn,
18143    x: &[f32],
18144    pool: Option<&Pool>,
18145) -> Vec<f32> {
18146    let t0 = std::time::Instant::now();
18147    let r = moe_ffn_route(m, x, pool, None);
18148    slot.note_route(t0.elapsed().as_nanos() as u64);
18149    match slot.forward(li, m, x, &r, pool) {
18150        Some(out) => out,
18151        None => crate::qtensor::float_activations_scope(|| {
18152            crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, &r, pool))
18153        }),
18154    }
18155}
18156
18157/// Verify rows share a bank frame; routing and fallback are decode's.
18158fn moe_ffn_banked_rows(
18159    slot: &mut crate::mimo_moe::Slot,
18160    li: usize,
18161    m: &MoeFfn,
18162    xs: &[f32],
18163    b: usize,
18164    hidden: usize,
18165    pool: Option<&Pool>,
18166) -> Vec<f32> {
18167    let t0 = std::time::Instant::now();
18168    let routes: Vec<_> = xs
18169        .chunks_exact(hidden)
18170        .map(|x| moe_ffn_route(m, x, pool, None))
18171        .collect();
18172    slot.note_route(t0.elapsed().as_nanos() as u64);
18173    if let Some(out) = slot.forward_rows(li, m, xs, &routes, pool) {
18174        return out;
18175    }
18176    let mut out = Vec::with_capacity(b * hidden);
18177    for (x, r) in xs.chunks_exact(hidden).zip(&routes) {
18178        let row = slot.forward(li, m, x, r, pool).unwrap_or_else(|| {
18179            // A failed bank must not stream missing experts into the arena.
18180            crate::qtensor::float_activations_scope(|| {
18181                crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, r, pool))
18182            })
18183        });
18184        out.extend(row);
18185    }
18186    out
18187}
18188
18189/// One-shot report of whether the whole-token wgpu graph actually formed.
18190/// A refusal silently reverts to the per-op path, which is how a model can
18191/// look "GPU-accelerated" while every layer walks the host.  A device prefix
18192/// is tracked separately because it still pays a host boundary for the tail.
18193fn graph_note(built: bool, layers_run: usize, total_layers: usize) {
18194    use std::sync::atomic::{AtomicBool, Ordering};
18195    if built {
18196        GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
18197        if total_layers > 0 && layers_run < total_layers {
18198            GRAPH_TOK_PREFIX.fetch_add(1, Ordering::Relaxed);
18199        } else {
18200            GRAPH_TOK_FULL.fetch_add(1, Ordering::Relaxed);
18201        }
18202    } else {
18203        GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
18204    }
18205    static SAID: AtomicBool = AtomicBool::new(false);
18206    if !SAID.swap(true, Ordering::Relaxed) {
18207        if built {
18208            tracing::info!("wgpu whole-token graph: ACTIVE");
18209        } else {
18210            tracing::warn!("wgpu whole-token graph refused — per-op path");
18211        }
18212    }
18213}
18214
18215/// Whole-token graph outcomes, process-wide: a benchmark that claims a
18216/// GPU number while MISS climbs is measuring the CPU — the honest-bench
18217/// contract makes that an error, not a footnote.
18218pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18219pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18220/// Graph calls that returned a hidden after running only a leading device
18221/// prefix.  These are valid hybrid executions but must not be reported as a
18222/// full GPU graph in benchmark evidence.
18223pub static GRAPH_TOK_PREFIX: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18224/// Graph calls that covered the complete requested layer span.
18225pub static GRAPH_TOK_FULL: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18226
18227/// Native Metal TokenGraph completion counters. These are incremented only
18228/// after checked command-buffer completion and successful readback, so a
18229/// fused-head NLL report can prove the route rather than infer it from env.
18230pub static METAL_GRAPH_TOK_OK: std::sync::atomic::AtomicU64 =
18231    std::sync::atomic::AtomicU64::new(0);
18232pub static METAL_GRAPH_HEAD_OK: std::sync::atomic::AtomicU64 =
18233    std::sync::atomic::AtomicU64::new(0);
18234pub static METAL_GRAPH_HEAD_MISS: std::sync::atomic::AtomicU64 =
18235    std::sync::atomic::AtomicU64::new(0);
18236pub static METAL_GRAPH_LAYERS: std::sync::atomic::AtomicU64 =
18237    std::sync::atomic::AtomicU64::new(0);
18238pub static METAL_GRAPH_ERRORS: std::sync::atomic::AtomicU64 =
18239    std::sync::atomic::AtomicU64::new(0);
18240/// Ordinary native-Metal rows-prefill admissions and completed rows.  These
18241/// counters are separate from TokenGraph token/head counts so a batch NLL
18242/// receipt cannot accidentally claim serial execution as batched.
18243pub static METAL_PREFILL_CHUNKS: std::sync::atomic::AtomicU64 =
18244    std::sync::atomic::AtomicU64::new(0);
18245pub static METAL_PREFILL_ROWS: std::sync::atomic::AtomicU64 =
18246    std::sync::atomic::AtomicU64::new(0);
18247pub static METAL_PREFILL_HEAD_ROWS: std::sync::atomic::AtomicU64 =
18248    std::sync::atomic::AtomicU64::new(0);
18249pub static METAL_PREFILL_ERRORS: std::sync::atomic::AtomicU64 =
18250    std::sync::atomic::AtomicU64::new(0);
18251
18252/// `CMF_MOE_BATCH=0` restores the per-expert serial loop — the A/B lever
18253/// for the batched kernel, and how its bit-identity is checked.
18254fn moe_batch_enabled() -> bool {
18255    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18256    *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
18257}
18258
18259/// Two-dispatch CPU MoE: every routed expert (and the shared one) fused
18260/// into one gate/up/SiLU dispatch and one down dispatch, instead of two
18261/// pool barriers per expert. Bit-identical to the serial loop below —
18262/// see `moe_gate_up_many` / `moe_down_many`. `None` = the batched kernel
18263/// does not cover this layer, walk the serial path.
18264fn moe_ffn_cpu_batched(
18265    m: &MoeFfn,
18266    x: &[f32],
18267    idx: &[usize],
18268    p: &[f32],
18269    wsum: f32,
18270    pool: Option<&Pool>,
18271) -> Option<Vec<f32>> {
18272    if idx.is_empty() || !moe_batch_enabled() {
18273        return None;
18274    }
18275    // The bake probe reads per-neuron activation mass out of the
18276    // single-expert path; batching would skip it. Rare and offline —
18277    // hand those runs to the serial loop.
18278    if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
18279        return None;
18280    }
18281    let n = idx.len() + usize::from(m.shared.is_some());
18282    let mut pairs = Vec::with_capacity(n);
18283    let mut downs = Vec::with_capacity(n);
18284    let mut ws = Vec::with_capacity(n);
18285    for &e in idx {
18286        let d = &m.experts[e];
18287        if d.act != Act::Silu {
18288            return None;
18289        }
18290        pairs.push((&d.gate_proj, &d.up_proj));
18291        downs.push(&d.down_proj);
18292        ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
18293    }
18294    // The shared expert goes last, matching the serial loop's order —
18295    // the f32 accumulation order is part of the bit-identity claim.
18296    if let Some((se, gate)) = &m.shared {
18297        if se.act != Act::Silu {
18298            return None;
18299        }
18300        let g = gate.as_ref().map_or(1.0, |gate| {
18301            let mut gl = [0.0f32; 1];
18302            gate.matvec(x, &mut gl, pool);
18303            1.0 / (1.0 + (-gl[0]).exp())
18304        });
18305        pairs.push((&se.gate_proj, &se.up_proj));
18306        downs.push(&se.down_proj);
18307        ws.push(g);
18308    }
18309    let inter = pairs[0].0.rows();
18310    let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
18311    if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
18312        return None;
18313    }
18314    let mut out = attention::take_buf(x.len());
18315    if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
18316        attention::recycle_buf(&mut out);
18317        return None;
18318    }
18319    Some(out)
18320}
18321
18322/// Exact CPU completion for the routed experts a dynamic device cache did
18323/// not contain. The weights are already the router's final normalized mix.
18324/// Keeping this independent of `MoeFfn` makes the job `Sync`: its routing
18325/// statistics live in a `RefCell`, while the immutable expert tensors can be
18326/// evaluated safely in parallel with the GPU's resident subset.
18327pub(crate) fn moe_cold_experts_cpu(
18328    experts: &[(&DenseFfn, f32)],
18329    x: &[f32],
18330    pool: Option<&Pool>,
18331) -> Vec<f32> {
18332    let mut out = attention::take_buf(x.len());
18333    if experts.is_empty() {
18334        return out;
18335    }
18336    let pairs: Vec<_> = experts
18337        .iter()
18338        .map(|(e, _)| (&e.gate_proj, &e.up_proj))
18339        .collect();
18340    let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
18341    let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
18342    let inter = experts[0].0.gate_proj.rows();
18343    let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
18344    if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
18345        && QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
18346    {
18347        return out;
18348    }
18349    out.fill(0.0);
18350    for &(expert, weight) in experts {
18351        let mut one = dense_ffn(expert, x, pool);
18352        for (o, v) in out.iter_mut().zip(&one) {
18353            *o += weight * v;
18354        }
18355        attention::recycle_buf(&mut one);
18356    }
18357    out
18358}
18359
18360/// Cold part of a short bank batch. Share each expert's weight stream
18361/// across its tokens, but reduce contributions in each token's route order.
18362/// On an unsupported CPU/layout, retain the single-token cold kernels.
18363pub(crate) fn moe_cold_experts_rows_cpu(
18364    jobs: &[Vec<(&DenseFfn, f32)>],
18365    xs: &[f32],
18366    hidden: usize,
18367    pool: Option<&Pool>,
18368) -> Vec<f32> {
18369    let mut out = vec![0.0; xs.len()];
18370    let mut experts: Vec<&DenseFfn> = Vec::new();
18371    let mut groups: Vec<Vec<usize>> = Vec::new();
18372    let mut terms = vec![Vec::new(); jobs.len()];
18373    for (r, row) in jobs.iter().enumerate() {
18374        for &(e, w) in row {
18375            let g = match experts.iter().position(|&d| std::ptr::eq(d, e)) {
18376                Some(g) => g,
18377                None => {
18378                    experts.push(e);
18379                    groups.push(Vec::new());
18380                    groups.len() - 1
18381                }
18382            };
18383            terms[r].push((g, groups[g].len(), w));
18384            groups[g].push(r);
18385        }
18386    }
18387    if experts.is_empty() {
18388        return out;
18389    }
18390    let pairs: Vec<_> = experts.iter().map(|e| (&e.gate_proj, &e.up_proj)).collect();
18391    let downs: Vec<_> = experts.iter().map(|e| &e.down_proj).collect();
18392    let lens: Vec<_> = groups.iter().map(Vec::len).collect();
18393    let count: usize = lens.iter().sum();
18394    let mut acts = vec![vec![0.0; experts[0].gate_proj.rows()]; count];
18395    let mut ds = vec![vec![0.0; hidden]; count];
18396    if QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut acts, pool)
18397        && QTensor::moe_down_rows(&downs, &lens, &acts, &mut ds, pool)
18398    {
18399        let mut offset = 0;
18400        let offsets: Vec<_> = lens
18401            .iter()
18402            .map(|&n| {
18403                let start = offset;
18404                offset += n;
18405                start
18406            })
18407            .collect();
18408        for (r, terms) in terms.iter().enumerate() {
18409            for &(g, slot, w) in terms {
18410                for (o, &v) in out[r * hidden..(r + 1) * hidden]
18411                    .iter_mut()
18412                    .zip(&ds[offsets[g] + slot])
18413                {
18414                    *o += w * v;
18415                }
18416            }
18417        }
18418    } else {
18419        for (r, jobs) in jobs.iter().enumerate() {
18420            if !jobs.is_empty() {
18421                let mut row = moe_cold_experts_cpu(jobs, &xs[r * hidden..(r + 1) * hidden], pool);
18422                out[r * hidden..(r + 1) * hidden].copy_from_slice(&row);
18423                attention::recycle_buf(&mut row);
18424            }
18425        }
18426    }
18427    out
18428}
18429
18430/// The pure-CPU MoE expert loop (also the fallback of every GPU refusal).
18431fn moe_ffn_cpu(
18432    m: &MoeFfn,
18433    x: &[f32],
18434    idx: &[usize],
18435    p: &[f32],
18436    wsum: f32,
18437    pool: Option<&Pool>,
18438) -> Vec<f32> {
18439    if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
18440        return out;
18441    }
18442    let mut out = attention::take_buf(x.len());
18443    for &e in idx {
18444        let mut eo = dense_ffn(&m.experts[e], x, pool);
18445        let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
18446        for i in 0..out.len() {
18447            out[i] += w * eo[i];
18448        }
18449        attention::recycle_buf(&mut eo);
18450    }
18451    if let Some((se, gate)) = &m.shared {
18452        let mut so = dense_ffn(se, x, pool);
18453        let g = gate.as_ref().map_or(1.0, |gate| {
18454            let mut gl = [0.0f32; 1];
18455            gate.matvec(x, &mut gl, pool);
18456            1.0 / (1.0 + (-gl[0]).exp())
18457        });
18458        for i in 0..out.len() {
18459            out[i] += g * so[i];
18460        }
18461        attention::recycle_buf(&mut so);
18462    }
18463    out
18464}
18465
18466/// DeepSeek-V2 MLA forward, expand-to-MHA form (see `AttnKind::Mla`):
18467/// per token the latent expands to every head's K/V and the ordinary
18468/// cache + grouped attend do the rest. K head layout is [rope | nope]
18469/// (rotary_dim = qk_rope rotates the shared rope key and each q head's
18470/// prefix); V rows are zero-padded to the K head_dim inside the cache
18471/// and the pad is sliced off before O. Attention importance is not
18472/// accumulated for MLA yet (no eviction interplay).
18473#[allow(clippy::too_many_arguments)]
18474pub(crate) fn mla_attention(
18475    w: &MlaWeights,
18476    normed: &[f32],
18477    cache: &mut crate::kv_cache::LayerKvCache,
18478    position: usize,
18479    inv_freq: &[f32],
18480    rope_scale: f32,
18481    eps: f64,
18482    pool: Option<&Pool>,
18483) -> Vec<f32> {
18484    let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
18485    let hd = dr + dn;
18486    let mut q = vec![0.0f32; nh * hd];
18487    match (&w.q_a, &w.q_a_norm) {
18488        (Some(qa), Some(qn)) => {
18489            let mut t = vec![0.0f32; qa.rows()];
18490            qa.matvec(normed, &mut t, pool);
18491            let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
18492            w.q_proj.matvec(&tn, &mut q, pool);
18493        }
18494        _ => w.q_proj.matvec(normed, &mut q, pool),
18495    }
18496    let mut ca = vec![0.0f32; lora + dr];
18497    w.kv_a.matvec(normed, &mut ca, pool);
18498    let (c_lat, k_rope) = ca.split_at_mut(lora);
18499    let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
18500    let mut kvb = vec![0.0f32; nh * (dn + dv)];
18501    w.kv_b.matvec(&latn, &mut kvb, pool);
18502    if !w.nope {
18503        attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
18504    }
18505    for h in 0..nh {
18506        if !w.nope {
18507            attention::rope_rotate_scaled(
18508                &mut q[h * hd..h * hd + dr],
18509                position,
18510                inv_freq,
18511                rope_scale,
18512            );
18513        }
18514    }
18515    let mut k = vec![0.0f32; nh * hd];
18516    let mut v = vec![0.0f32; nh * hd];
18517    for h in 0..nh {
18518        k[h * hd..h * hd + dr].copy_from_slice(k_rope);
18519        k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
18520        v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
18521    }
18522    cache.append(&k, &v, &vec![true; nh]);
18523    let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
18524    attention::recycle_buf(&mut imp);
18525    let mut ov = vec![0.0f32; nh * dv];
18526    for h in 0..nh {
18527        ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
18528    }
18529    let mut out = vec![0.0f32; w.o_proj.rows()];
18530    w.o_proj.matvec(&ov, &mut out, pool);
18531    out
18532}
18533
18534/// Gemma-4 dual-branch FFN (spec: see `FfnKind::DenseMoe`). The dense
18535/// branch reads the pre-FFN-normed activation; the router and the
18536/// expert branch read the RAW residual — the router through a
18537/// scale-less rms norm (its constant gain is folded into the weights),
18538/// the experts through `pre_norm_2`. CPU path; GPU graphs refuse the
18539/// layer kind honestly.
18540fn dense_moe_ffn(
18541    dm: &DenseMoeFfn,
18542    x_normed: &[f32],
18543    h_raw: &[f32],
18544    eps: f64,
18545    norm_style: NormStyle,
18546    pool: Option<&Pool>,
18547) -> Vec<f32> {
18548    let mut d = dense_ffn(&dm.dense, x_normed, pool);
18549    d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
18550    let m = &dm.moe;
18551    let ne = m.experts.len();
18552    let mut logits = vec![0.0f32; ne];
18553    if m.router_input_norm {
18554        let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
18555        let inv = 1.0 / (ss + eps as f32).sqrt();
18556        let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
18557        m.router.matvec(&xr, &mut logits, pool);
18558    } else {
18559        m.router.matvec(h_raw, &mut logits, pool);
18560    }
18561    let (idx, p, wsum) = moe_route(&logits, m, None);
18562    {
18563        let mut st = m.stats.borrow_mut();
18564        if st.len() < ne {
18565            st.resize(ne, 0);
18566        }
18567        for &e in &idx {
18568            st[e] += 1;
18569        }
18570    }
18571    let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
18572    let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
18573    let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
18574    for (di, mi) in d.iter_mut().zip(&mo) {
18575        *di += mi;
18576    }
18577    d
18578}
18579
18580/// Building the MoE-layer GPU jobs: all selected experts (+shared) must
18581/// be q8_2f-Mapped from the primary mapping; otherwise None → CPU path.
18582/// One-shot report of why the MoE GPU block refused. A silent `?` here
18583/// sends every expert to the CPU with nothing in the logs to say so —
18584/// which is exactly how a q4tp MoE model looked "GPU-accelerated" while
18585/// running entirely on the host.
18586fn moe_gpu_refused(why: &'static str) {
18587    use std::sync::atomic::{AtomicBool, Ordering};
18588    static SAID: AtomicBool = AtomicBool::new(false);
18589    if !SAID.swap(true, Ordering::Relaxed) {
18590        tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
18591    }
18592}
18593
18594fn moe_ffn_gpu(
18595    m: &MoeFfn,
18596    x: &[f32],
18597    idx: &[usize],
18598    p: &[f32],
18599    wsum: f32,
18600    pool: Option<&Pool>,
18601) -> Option<Vec<f32>> {
18602    use crate::gpu::MoeJob;
18603
18604    let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
18605    let mut model_ref = None;
18606    for &e in idx {
18607        if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
18608            moe_gpu_refused("push_job(expert)");
18609            return None;
18610        }
18611    }
18612    if let Some((se, gate)) = &m.shared {
18613        let g = gate.as_ref().map_or(1.0, |gate| {
18614            let mut gl = [0.0f32; 1];
18615            gate.matvec(x, &mut gl, pool);
18616            1.0 / (1.0 + (-gl[0]).exp())
18617        });
18618        if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
18619            moe_gpu_refused("push_job(shared)");
18620            return None;
18621        }
18622    }
18623    let Some(model) = model_ref else {
18624        moe_gpu_refused("no model_ref");
18625        return None;
18626    };
18627    let hidden = jobs[0].down.1;
18628    let mut out = vec![0.0f32; hidden];
18629    if crate::gpu::moe_block(&model, &jobs, &mut out) {
18630        Some(out)
18631    } else {
18632        moe_gpu_refused("gpu::moe_block");
18633        None
18634    }
18635}
18636
18637/// Single-position FFN dispatch.
18638fn ffn_forward(
18639    ffn: &FfnKind,
18640    x: &[f32],
18641    pool: Option<&Pool>,
18642    experts_allowed: Option<&[bool]>,
18643) -> Vec<f32> {
18644    match ffn {
18645        FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
18646        FfnKind::Dense(d) => dense_ffn(d, x, pool),
18647        FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
18648        // Dual-branch layers need the raw residual — their callers
18649        // dispatch dense_moe_ffn directly; the auxiliary paths that land
18650        // here (MTP draft, o1 replay) do not co-occur with gemma-4 MoE.
18651        FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18652    }
18653}
18654
18655/// Fused two-position FFN: gate/up/down streamed once (dense). MoE
18656/// falls back to two singles — expert sets differ per position, there
18657/// is nothing to fuse.
18658fn ffn_forward_pair(
18659    ffn: &FfnKind,
18660    x1: &[f32],
18661    x2: &[f32],
18662    pool: Option<&Pool>,
18663    experts_allowed: Option<&[bool]>,
18664) -> (Vec<f32>, Vec<f32>) {
18665    let d = match ffn {
18666        // A tube layer has nothing to fuse across the pair — the tubes
18667        // are separate matrices; two singles are the honest path.
18668        FfnKind::Dense(d) if !d.segs.is_empty() => {
18669            return (
18670                tube_ffn(d, x1, 1, pool, None),
18671                tube_ffn(d, x2, 1, pool, None),
18672            );
18673        }
18674        FfnKind::Dense(d) => d,
18675        FfnKind::Moe(m) => {
18676            return (
18677                moe_ffn(m, x1, pool, experts_allowed),
18678                moe_ffn(m, x2, pool, experts_allowed),
18679            );
18680        }
18681        FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18682    };
18683    let inter = d.gate_proj.rows();
18684    FFN_SCRATCH.with(|s| {
18685        let mut s = s.borrow_mut();
18686        let [g1, g2, u1, u2] = &mut *s;
18687        g1.resize(inter, 0.0);
18688        g2.resize(inter, 0.0);
18689        u1.resize(inter, 0.0);
18690        u2.resize(inter, 0.0);
18691        // Multi-matrix pair job: gate+up under one pool dispatch
18692        // (o1s = lane-1 outputs across tensors, o2s = lane-2).
18693        QTensor::matvec2_many(
18694            [&d.gate_proj, &d.up_proj],
18695            x1,
18696            x2,
18697            [g1.as_mut_slice(), u1.as_mut_slice()],
18698            [g2.as_mut_slice(), u2.as_mut_slice()],
18699            pool,
18700        );
18701        for i in 0..inter {
18702            g1[i] = d.act.combine(g1[i], u1[i]);
18703            g2[i] = d.act.combine(g2[i], u2[i]);
18704        }
18705        let mut o1 = attention::take_buf(d.down_proj.rows());
18706        let mut o2 = attention::take_buf(d.down_proj.rows());
18707        d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
18708        (o1, o2)
18709    })
18710}
18711
18712#[cfg(test)]
18713mod tests {
18714
18715    /// The 0.7.6 prefill-chunk rule: a plain dense stack wholly on a
18716    /// discrete card reads the prompt in wide chunks on x86; every other
18717    /// case keeps the width it had (the GDN-hybrid, MoE and DeepSeek paths
18718    /// were tuned on hardware not measured for this change).
18719    #[test]
18720    fn prefill_chunk_rule_widens_only_dense_on_discrete() {
18721        use super::{
18722            prefill_chunk_rule, ChunkHost, ChunkStackFacts, DISCRETE_DENSE_PREFILL_CHUNK,
18723        };
18724        let dense_card = ChunkStackFacts {
18725            plain_dense: true,
18726            discrete: true,
18727            gpu_on: true,
18728            ..Default::default()
18729        };
18730        assert!(dense_card.dense_on_discrete());
18731        // The bug: a dense Llama on a Vulkan RTX 3090 got 48.
18732        assert_eq!(
18733            prefill_chunk_rule(None, ChunkHost::Other, dense_card.dense_on_discrete()),
18734            DISCRETE_DENSE_PREFILL_CHUNK
18735        );
18736        assert!(DISCRETE_DENSE_PREFILL_CHUNK > 48);
18737        for (label, facts) in [
18738            ("GDN hybrid / MoE / DeepSeek stack", ChunkStackFacts { plain_dense: false, ..dense_card }),
18739            ("integrated GPU", ChunkStackFacts { discrete: false, ..dense_card }),
18740            ("CPU only", ChunkStackFacts { gpu_on: false, discrete: false, ..dense_card }),
18741            ("capacity split", ChunkStackFacts { capacity_split: true, ..dense_card }),
18742            ("multi-GPU plan", ChunkStackFacts { multi_gpu: true, ..dense_card }),
18743            ("O(1) layers", ChunkStackFacts { o1: true, ..dense_card }),
18744        ] {
18745            assert!(!facts.dense_on_discrete(), "{label}");
18746            assert_eq!(
18747                prefill_chunk_rule(None, ChunkHost::Other, facts.dense_on_discrete()),
18748                48,
18749                "{label} keeps the historical x86 chunk"
18750            );
18751        }
18752        // Other hosts are untouched whatever the model.
18753        for dense in [false, true] {
18754            assert_eq!(prefill_chunk_rule(None, ChunkHost::Macos, dense), 512);
18755            assert_eq!(prefill_chunk_rule(None, ChunkHost::Aarch64, dense), 256);
18756        }
18757        // CMF_PREFILL_CHUNK still wins everywhere (and is clamped to ≥ 1).
18758        for host in [ChunkHost::Macos, ChunkHost::Aarch64, ChunkHost::Other] {
18759            for dense in [false, true] {
18760                assert_eq!(prefill_chunk_rule(Some(48), host, dense), 48);
18761                assert_eq!(prefill_chunk_rule(Some(0), host, dense), 1);
18762            }
18763        }
18764    }
18765
18766    #[test]
18767    fn kv_reuse_plan_pulls_rows_decode_wrote_only_on_the_device() {
18768        use super::{ReuseLayer, ReusePlan, kv_reuse_plan};
18769        let full = |host_rows, device_rows| ReuseLayer {
18770            full: true,
18771            host_rows,
18772            device_rows,
18773            device_state: false,
18774        };
18775        // Turn 1: 300-token prompt prefilled on the host, 40 tokens decoded
18776        // by the wgpu graph into the device mirror only. Turn 2 reuses 339.
18777        assert_eq!(
18778            kv_reuse_plan(339, &[full(300, Some(339)), full(300, Some(339))]),
18779            ReusePlan::Pull(vec![(0, 300, 339), (1, 300, 339)])
18780        );
18781        // CPU / Metal: the host owner already holds every forwarded row.
18782        assert_eq!(kv_reuse_plan(339, &[full(339, None)]), ReusePlan::Ready);
18783        // A mirror past the prefix is fine for the host (it gets rewound).
18784        assert_eq!(kv_reuse_plan(339, &[full(339, Some(345))]), ReusePlan::Ready);
18785        // GPU prefix / CPU tail: only the device layers lag.
18786        assert_eq!(
18787            kv_reuse_plan(339, &[full(300, Some(339)), full(339, None)]),
18788            ReusePlan::Pull(vec![(0, 300, 339)])
18789        );
18790        // The device cannot supply the missing rows: never continue.
18791        assert_eq!(kv_reuse_plan(339, &[full(300, Some(320))]), ReusePlan::Fresh);
18792        assert_eq!(kv_reuse_plan(339, &[full(300, None)]), ReusePlan::Fresh);
18793        assert_eq!(kv_reuse_plan(339, &[full(350, None)]), ReusePlan::Fresh);
18794        // A recurrent state advanced on the device cannot be handed to a
18795        // host prefill (it is not rewindable and the host copy is stale).
18796        let conv = |device_state| ReuseLayer {
18797            full: false,
18798            host_rows: 0,
18799            device_rows: None,
18800            device_state,
18801        };
18802        assert_eq!(
18803            kv_reuse_plan(339, &[conv(true), full(300, Some(339))]),
18804            ReusePlan::Fresh
18805        );
18806        assert_eq!(kv_reuse_plan(339, &[conv(false), full(339, None)]), ReusePlan::Ready);
18807    }
18808
18809    #[test]
18810    fn nll_graph_policy_scopes_only_the_fused_head() {
18811        for (label, unmasked, prefer_graph, native_metal, want_graph, want_head) in [
18812            // A Vulkan/Wgpu hidden-only graph remains the quality route.
18813            ("vulkan graph", true, true, false, true, false),
18814            // Native Metal adds the strict fused graph-head contract.
18815            ("native Metal graph", true, true, true, true, true),
18816            // Masked NLL and the explicit non-graph fallback remain unchanged.
18817            ("masked", false, true, false, false, false),
18818            ("graph disabled", true, false, true, false, false),
18819        ] {
18820            let (graph_quality, graph_head_required) =
18821                super::nll_graph_policy(unmasked, prefer_graph, native_metal);
18822            assert_eq!(graph_quality, want_graph, "{label}: graph quality");
18823            assert_eq!(graph_head_required, want_head, "{label}: fused head");
18824        }
18825    }
18826
18827    #[test]
18828    fn mtp_prefill_pair_boundaries_skip_only_final_prompt_row() {
18829        assert_eq!(mtp_prefill_pair_count(0, 128, 256), 128);
18830        assert_eq!(mtp_prefill_pair_count(128, 256, 256), 127);
18831        assert_eq!(mtp_prefill_pair_count(0, 256, 256), 255);
18832        assert_eq!(mtp_prefill_pair_count(256, 256, 256), 0);
18833        assert_eq!(mtp_prefill_pair_count(300, 320, 256), 0);
18834    }
18835
18836    #[test]
18837    fn cancel_flag_stops_generation() {
18838        let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
18839        // Set before the call: the prefill loops honour it, the run
18840        // returns immediately with the cancelled reason and no tokens.
18841        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
18842        let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
18843        assert_eq!(r.finish_reason, "cancelled");
18844        assert!(
18845            r.token_ids.is_empty(),
18846            "no tokens after cancel: {:?}",
18847            r.token_ids
18848        );
18849        assert_eq!(p.kv_cache.seq_len(), 0);
18850        assert!(p.kv_history.is_empty());
18851        assert!(!p.graph_want_logits);
18852        assert!(p.graph_logits.is_none());
18853        // Flag auto-cleared: the next call generates normally.
18854        let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
18855        assert_ne!(r2.finish_reason, "cancelled");
18856    }
18857    use super::*;
18858
18859    /// sparse_ffn_quant must equal a dense FFN where inactive neurons are
18860    /// zeroed (mask × mmap correctness). On F32 tensors this is EXACT —
18861    /// it validates the row_dot / add_col_scaled / scatter indexing, the
18862    /// bug-prone part. The q8 branches reuse the golden-tested linear
18863    /// The per-token sparse path reads a transposed `down`; it must
18864    /// agree with the arm that computes everything and zeroes the
18865    /// losers, or the speed measurement is measuring a different model.
18866    #[test]
18867    fn dynamic_ffn_equals_the_zeroing_arm() {
18868        let (hidden, inter) = (8usize, 32usize);
18869        let synth = |n: usize, salt: usize| -> Vec<f32> {
18870            (0..n)
18871                .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
18872                .collect()
18873        };
18874        let down = synth(hidden * inter, 3);
18875        let mut down_t = vec![0.0f32; inter * hidden];
18876        for r in 0..hidden {
18877            for c in 0..inter {
18878                down_t[c * hidden + r] = down[r * inter + c];
18879            }
18880        }
18881        let d = DenseFfn {
18882            gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18883            up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18884            down_proj: QTensor::from_f32(down.clone(), hidden, inter),
18885            act: Act::Silu,
18886            down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
18887            segs: Vec::new(),
18888        };
18889        let x = synth(hidden, 11);
18890        let k = 12usize;
18891        let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
18892        // Reference: full compute, keep the k loudest |silu(gate)|.
18893        let mut g = vec![0.0f32; inter];
18894        d.gate_proj.matvec(&x, &mut g, None);
18895        let mut u = vec![0.0f32; inter];
18896        d.up_proj.matvec(&x, &mut u, None);
18897        for v in g.iter_mut() {
18898            *v = inference::silu(*v);
18899        }
18900        keep_top_k(&mut g, k);
18901        for i in 0..inter {
18902            g[i] *= u[i];
18903        }
18904        let mut want = vec![0.0f32; hidden];
18905        d.down_proj.matvec(&g, &mut want, None);
18906        for (a, b) in want.iter().zip(&got) {
18907            assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
18908        }
18909    }
18910
18911    /// A tube layer is the same layer, re-cut. With every tube open the
18912    /// answer must equal the dense FFN over the concatenated neurons
18913    /// (the permutation is an identity on the layer's function); with a
18914    /// tube closed it must equal the dense FFN with those neurons
18915    /// zeroed — the mask semantics, now paid for in bytes not read.
18916    #[test]
18917    fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
18918        let (hidden, core, tube) = (8usize, 12usize, 8usize);
18919        let inter = core + tube;
18920        let synth = |n: usize, salt: usize| -> Vec<f32> {
18921            (0..n)
18922                .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
18923                .collect()
18924        };
18925        let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
18926        let d_all = synth(hidden * inter, 3);
18927        // The dense layer, and the same weights cut into core + tube.
18928        let dense = DenseFfn {
18929            gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
18930            up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
18931            down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
18932            act: Act::Silu,
18933            down_t: None,
18934            segs: Vec::new(),
18935        };
18936        let rows =
18937            |v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
18938        let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
18939            let mut o = Vec::with_capacity(hidden * (b - a));
18940            for r in 0..hidden {
18941                o.extend_from_slice(&v[r * inter + a..r * inter + b]);
18942            }
18943            o
18944        };
18945        let tubed = DenseFfn {
18946            down_t: None,
18947            gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
18948            up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
18949            down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
18950            act: Act::Silu,
18951            segs: vec![FfnSeg {
18952                gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
18953                up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
18954                down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
18955                start: core,
18956                width: tube,
18957            }],
18958        };
18959        let x = synth(hidden, 7);
18960        let want = dense_ffn(&dense, &x, None);
18961        let got = tube_ffn(&tubed, &x, 1, None, None);
18962        for (a, b) in want.iter().zip(&got) {
18963            assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
18964        }
18965        // Closed tube: bits on for the core, off for the tube.
18966        let mut bits = vec![0u8; inter.div_ceil(8)];
18967        for n in 0..core {
18968            bits[n / 8] |= 1 << (n % 8);
18969        }
18970        let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18971        let masked = dense_ffn_masked(&dense, &x, None, &bits);
18972        for (a, b) in masked.iter().zip(&closed) {
18973            assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
18974        }
18975        // The batched arm must agree with the single-position one.
18976        let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18977        for (a, b) in closed.iter().zip(&batch) {
18978            assert_eq!(a, b, "batch arm disagrees with decode arm");
18979        }
18980    }
18981
18982    /// scale, structurally identical to the matvec kernels.
18983    #[test]
18984    fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
18985        let (hidden, inter) = (16usize, 40usize);
18986        let synth = |n: usize, salt: usize| -> Vec<f32> {
18987            (0..n)
18988                .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
18989                .collect()
18990        };
18991        let d = DenseFfn {
18992            gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18993            up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18994            down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
18995            act: Act::Silu,
18996            down_t: None,
18997            segs: Vec::new(),
18998        };
18999        let x = synth(hidden, 9);
19000        // Active = every 3rd neuron.
19001        let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
19002
19003        let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
19004
19005        // Reference: full dense FFN but g[i]=0 for inactive neurons.
19006        let mut g = vec![0.0f32; inter];
19007        d.gate_proj.matvec(&x, &mut g, None);
19008        let mut u = vec![0.0f32; inter];
19009        d.up_proj.matvec(&x, &mut u, None);
19010        let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
19011        for i in 0..inter {
19012            g[i] = if act_set.contains(&(i as u16)) {
19013                inference::silu(g[i]) * u[i]
19014            } else {
19015                0.0
19016            };
19017        }
19018        let mut reference = vec![0.0f32; hidden];
19019        d.down_proj.matvec(&g, &mut reference, None);
19020
19021        let max_d = sparse
19022            .iter()
19023            .zip(&reference)
19024            .map(|(a, b)| (a - b).abs())
19025            .fold(0.0f32, f32::max);
19026        assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
19027    }
19028
19029    /// Attach a synthetic MTP head (same structure as a main layer).
19030    fn attach_test_mtp(p: &mut Pipeline) {
19031        let (h, inter, heads, kv, hd) = (
19032            p.hidden_size,
19033            p.intermediate_size,
19034            p.num_heads,
19035            p.num_kv_heads,
19036            p.head_dim,
19037        );
19038        let synth = |n: usize, salt: usize| -> Vec<f32> {
19039            (0..n)
19040                .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
19041                .collect()
19042        };
19043        let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
19044            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19045        };
19046        p.mtp = Some(MtpModule {
19047            enorm: vec![1.0; h],
19048            hnorm: vec![1.0; h],
19049            eh_proj: qt(h, 2 * h, 301),
19050            layer: LayerWeights {
19051                input_norm: vec![1.0; h],
19052                post_norm: vec![1.0; h],
19053                attn_out_norm: None,
19054                ffn_out_norm: None,
19055                layer_scale: None,
19056                ffn: FfnKind::Dense(DenseFfn {
19057                    gate_proj: qt(inter, h, 315),
19058                    up_proj: qt(inter, h, 316),
19059                    down_proj: qt(h, inter, 317),
19060                    act: Act::Silu,
19061                    down_t: None,
19062                    segs: Vec::new(),
19063                }),
19064                attn: AttnKind::Full {
19065                    bias: None,
19066                    wq: qt(heads * hd, h, 311),
19067                    wk: qt(kv * hd, h, 312),
19068                    wv: qt(kv * hd, h, 313),
19069                    wo: qt(h, heads * hd, 314),
19070                    q_norm: None,
19071                    k_norm: None,
19072                    output_gate: false,
19073                    softplus_gate: None,
19074                },
19075            },
19076            final_norm: vec![1.0; h],
19077            kv: crate::kv_cache::LayerKvCache::new(kv, hd),
19078        });
19079    }
19080
19081    #[test]
19082    fn speculative_equals_vanilla_greedy() {
19083        // Speculative decode and the wgpu token graph are mutually
19084        // exclusive; a leaked CMF_GPU=wgpu from a parallel gpu test
19085        // would silently disable drafting. Pin the graph off.
19086        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19087        let run = |spec: bool| {
19088            let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19089            p.sampler_config.temperature = 0.0;
19090            attach_test_mtp(&mut p);
19091            p.speculative = spec;
19092            let r = p.generate("abcdef", 12, None, None).unwrap();
19093            (r.token_ids, r.mtp_drafted, r.mtp_accepted)
19094        };
19095        let (vanilla, d0, _) = run(false);
19096        let (spec, d1, a1) = run(true);
19097        assert_eq!(d0, 0, "vanilla path must not draft");
19098        assert!(d1 > 0, "speculative path must draft");
19099        assert_eq!(
19100            vanilla, spec,
19101            "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
19102        );
19103    }
19104
19105    #[test]
19106    fn speculative_accepts_constant_oracle() {
19107        // See speculative_equals_vanilla_greedy: pin the wgpu graph off.
19108        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19109        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19110        p.sampler_config.temperature = 0.0;
19111        p.sampler_config.repetition_penalty = 1.0;
19112        // Constant lm_head → every logit equal → both the main model and
19113        // the draft head argmax to token 0: acceptance must be 100%.
19114        p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
19115        attach_test_mtp(&mut p);
19116        p.speculative = true;
19117        let r = p.generate("abcd", 10, None, None).unwrap();
19118        assert!(r.mtp_drafted > 0);
19119        assert_eq!(
19120            r.mtp_accepted, r.mtp_drafted,
19121            "constant logits → every draft accepted"
19122        );
19123        // Ties resolve to the same token in both the main and draft
19124        // heads — the sequence is one repeated token.
19125        assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
19126    }
19127
19128    #[test]
19129    fn empty_prompt_is_an_error_not_a_panic() {
19130        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
19131        let r = p.generate("", 4, None, None);
19132        assert!(r.is_err(), "empty prompt must be a clean error");
19133    }
19134
19135    #[test]
19136    fn every_token_enters_kv_exactly_once() {
19137        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19138        // Greedy so no RNG variance; byte tokenizer → 3 prompt tokens.
19139        p.sampler_config.temperature = 0.0;
19140        let r = p.generate("abc", 2, None, None).unwrap();
19141        assert_eq!(r.prompt_tokens, 3);
19142        // prompt(3) + first sampled token forwarded before second logits:
19143        // step0 samples from prefill hidden (no extra forward), then
19144        // forwards t1 → cache 4; step1 samples, loop ends (max_tokens).
19145        assert_eq!(
19146            p.kv_cache.seq_len(),
19147            3 + r.tokens_generated - 1,
19148            "each token must be cached exactly once (v1 cached the last prompt token twice)"
19149        );
19150    }
19151
19152    #[test]
19153    fn generation_is_reproducible_with_seed() {
19154        let run = || {
19155            let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19156            p.generate("hello", 8, None, None).unwrap().token_ids
19157        };
19158        assert_eq!(run(), run());
19159    }
19160
19161    #[test]
19162    fn resetting_sampler_restarts_the_seeded_stream() {
19163        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19164        let config = SamplerConfig {
19165            seed: Some(1234),
19166            ..SamplerConfig::default()
19167        };
19168        p.set_sampler_config(config.clone());
19169        let first = p.generate("hello", 8, None, None).unwrap().token_ids;
19170        p.set_sampler_config(config);
19171        let second = p.generate("hello", 8, None, None).unwrap().token_ids;
19172        assert_eq!(first, second);
19173    }
19174
19175    #[test]
19176    fn eviction_bounds_the_cache() {
19177        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
19178        p.kv_cache.max_seq_len = 6;
19179        p.sampler_config.temperature = 0.0;
19180        let _ = p.generate("abcd", 12, None, None).unwrap();
19181        assert!(
19182            p.kv_cache.seq_len() <= 6 + 1,
19183            "cache must stay bounded by max_seq_len (got {})",
19184            p.kv_cache.seq_len()
19185        );
19186    }
19187
19188    #[test]
19189    fn confidence_matches_tokens_and_is_a_probability() {
19190        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19191        p.sampler_config.temperature = 0.0;
19192        p.sampler_config.repetition_penalty = 1.0;
19193        let r = p.generate("abcd", 10, None, None).unwrap();
19194        assert_eq!(
19195            r.token_confidence.len(),
19196            r.token_ids.len(),
19197            "one confidence per emitted token"
19198        );
19199        for &c in &r.token_confidence {
19200            assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
19201        }
19202        // top1_prob is a valid softmax probability.
19203        let logits = [1.0f32, 3.0, 0.5, 3.0];
19204        let p0 = top1_prob_t(&logits, 1, 1.0);
19205        let p1 = top1_prob_t(&logits, 3, 1.0);
19206        assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
19207        assert!(p0 > 0.0 && p0 < 1.0);
19208        // Calibration temperature > 1 softens an over-confident peak.
19209        let sharp = top1_prob_t(&logits, 1, 1.0);
19210        let soft = top1_prob_t(&logits, 1, 2.0);
19211        assert!(soft < sharp, "higher temperature lowers peak confidence");
19212    }
19213
19214    #[test]
19215    fn trace_is_opt_in_and_parallels_the_output() {
19216        // Off by default: the runtime is silent unless observation asked.
19217        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19218        p.sampler_config.temperature = 0.0;
19219        p.sampler_config.repetition_penalty = 1.0;
19220        let r = p.generate("abcd", 10, None, None).unwrap();
19221        assert!(r.traces.is_empty(), "trace must be empty unless enabled");
19222
19223        // On: exactly one row per emitted token, aligned with the output.
19224        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19225        p.sampler_config.temperature = 0.0;
19226        p.sampler_config.repetition_penalty = 1.0;
19227        p.set_trace(true);
19228        let r = p.generate("abcd", 10, None, None).unwrap();
19229        assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
19230        for (i, tr) in r.traces.iter().enumerate() {
19231            assert_eq!(tr.t, i, "trace index is sequential");
19232            assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
19233            assert_eq!(
19234                tr.confidence, r.token_confidence[i],
19235                "trace confidence matches the confidence channel"
19236            );
19237            // No dynamic router in this pipeline → no skill, no coherence.
19238            assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
19239        }
19240    }
19241
19242    #[test]
19243    fn explain_prefill_logits_match_greedy_first_token() {
19244        // `cortiq explain` shows the next-token distribution from
19245        // prefill_next_logits; its argmax must equal what greedy generate
19246        // actually emits first — otherwise explain would lie.
19247        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19248        p.sampler_config.temperature = 0.0;
19249        p.sampler_config.repetition_penalty = 1.0;
19250        let ids = p.tokenizer.encode("abcd");
19251        let logits = p.prefill_next_logits(&ids, None);
19252        let argmax = logits
19253            .iter()
19254            .enumerate()
19255            .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
19256            .unwrap()
19257            .0 as u32;
19258        let r = p.generate("abcd", 1, None, None).unwrap();
19259        assert_eq!(
19260            argmax, r.token_ids[0],
19261            "explain preview must match greedy emit"
19262        );
19263    }
19264
19265    #[test]
19266    fn laguna_shared_expert_is_unconditionally_added() {
19267        let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
19268        let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
19269        let zero_dense = || DenseFfn {
19270            gate_proj: matrix(vec![0.0; 4]),
19271            up_proj: matrix(vec![0.0; 4]),
19272            down_proj: matrix(vec![0.0; 4]),
19273            act: Act::Silu,
19274            down_t: None,
19275            segs: Vec::new(),
19276        };
19277        let shared = DenseFfn {
19278            gate_proj: identity(),
19279            up_proj: identity(),
19280            down_proj: identity(),
19281            act: Act::Silu,
19282            down_t: None,
19283            segs: Vec::new(),
19284        };
19285        let x = [1.0, 2.0];
19286        let expected = dense_ffn(&shared, &x, None);
19287        let moe = MoeFfn {
19288            router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
19289            experts: vec![zero_dense()],
19290            top_k: 1,
19291            norm_topk_prob: true,
19292            router_sigmoid: true,
19293            expert_bias: None,
19294            routed_scaling: 1.0,
19295            route_tau: None,
19296            shared: Some((shared, None)),
19297            stats: std::cell::RefCell::new(Vec::new()),
19298            act_sq: std::cell::RefCell::new(Vec::new()),
19299            act_rows: std::cell::RefCell::new(Vec::new()),
19300            mask: None,
19301            per_expert_scale: None,
19302            router_input_norm: false,
19303            resonance: None,
19304            grown: Vec::new(),
19305        };
19306        let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
19307        for (actual, expected) in actual.iter().zip(expected) {
19308            assert!((actual - expected).abs() < 1e-6);
19309        }
19310    }
19311
19312    /// A tiny MiMo-V2-shaped stack (the M3 fixture): layers [full, sliding,
19313    /// sliding, full]; 4 Q heads over 1 (full) / 2 (sliding) KV heads;
19314    /// head_dim 8 with 4-wide V heads; partial rotary 4 at θ 1e7 (full) /
19315    /// 1e4 (sliding); window 3; learned sinks on the sliding layers; layer
19316    /// 0 a dense FFN, layers 1..3 sigmoid-routed MoE with a selection bias
19317    /// (4 experts, top-2, renormalized, no shared expert). Geometry and
19318    /// sinks go through the same `set_attn_geometry` / `set_layer_sinks`
19319    /// the loader calls.
19320    fn mimo_test_pipeline() -> Pipeline {
19321        let (hs, inter, nh, hd, vd, vocab) = (16usize, 24usize, 4usize, 8usize, 4usize, 64usize);
19322        let kvh = [1usize, 2, 2, 1];
19323        let synth = |n: usize, salt: usize| -> Vec<f32> {
19324            (0..n)
19325                .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
19326                .collect()
19327        };
19328        let qt = |rows: usize, cols: usize, salt: usize| {
19329            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19330        };
19331        let dense = |inter: usize, salt: usize| DenseFfn {
19332            gate_proj: qt(inter, hs, salt),
19333            up_proj: qt(inter, hs, salt + 1),
19334            down_proj: qt(hs, inter, salt + 2),
19335            act: Act::Silu,
19336            down_t: None,
19337            segs: Vec::new(),
19338        };
19339        let layers: Vec<LayerWeights> = (0..4)
19340            .map(|li| LayerWeights {
19341                input_norm: vec![1.0; hs],
19342                post_norm: vec![1.0; hs],
19343                attn_out_norm: None,
19344                ffn_out_norm: None,
19345                layer_scale: None,
19346                ffn: if li == 0 {
19347                    FfnKind::Dense(dense(inter, 50))
19348                } else {
19349                    FfnKind::Moe(MoeFfn {
19350                        router: qt(4, hs, 60 + li),
19351                        experts: (0..4).map(|e| dense(8, 70 + li * 10 + e * 3)).collect(),
19352                        top_k: 2,
19353                        norm_topk_prob: true,
19354                        router_sigmoid: true,
19355                        expert_bias: Some(vec![0.02, -0.03, 0.01, 0.0]),
19356                        routed_scaling: 1.0,
19357                        route_tau: None,
19358                        shared: None,
19359                        stats: std::cell::RefCell::new(Vec::new()),
19360                        act_sq: std::cell::RefCell::new(Vec::new()),
19361                        act_rows: std::cell::RefCell::new(Vec::new()),
19362                        mask: None,
19363                        per_expert_scale: None,
19364                        router_input_norm: false,
19365                        resonance: None,
19366                        grown: Vec::new(),
19367                    })
19368                },
19369                attn: AttnKind::Full {
19370                    wq: qt(nh * hd, hs, li * 10 + 1),
19371                    wk: qt(kvh[li] * hd, hs, li * 10 + 2),
19372                    wv: qt(kvh[li] * vd, hs, li * 10 + 3),
19373                    wo: qt(hs, nh * vd, li * 10 + 4),
19374                    q_norm: None,
19375                    k_norm: None,
19376                    output_gate: false,
19377                    softplus_gate: None,
19378                    bias: None,
19379                },
19380            })
19381            .collect();
19382        let mut p = Pipeline::new(
19383            Tokenizer::byte_level(),
19384            PipelineWeights {
19385                embed_tokens: qt(vocab, hs, 100),
19386                layers,
19387                lm_head: qt(vocab, hs, 200),
19388                final_norm: vec![1.0; hs],
19389            },
19390            hs,
19391            inter,
19392            nh,
19393            1, // header num_kv_heads (the full layers')
19394            hd,
19395            4,
19396            4,
19397            false,
19398            vocab,
19399            1e-6,
19400            1e7,
19401            NormStyle::Qwen,
19402            4096,
19403            SamplerConfig {
19404                seed: Some(7),
19405                ..Default::default()
19406            },
19407        );
19408        // Diagnostics stay off whatever the test environment exports.
19409        p.layer_dump = None;
19410        p.set_rotary(4, 1e7);
19411        p.sliding_layers = Some(vec![false, true, true, false]);
19412        p.swa = Some((3, usize::MAX));
19413        p.rotary_dim_local = Some(4);
19414        p.inv_freq_local = Some(std::sync::Arc::new(attention::rope_inv_freq(4, 1e4)));
19415        p.set_attn_geometry(Some(kvh.to_vec()), Some(vd)).unwrap();
19416        p.set_layer_sinks(1, vec![0.5, -1.0, 1.5, 0.0]).unwrap();
19417        p.set_layer_sinks(2, vec![-0.25, 2.0, 0.75, -1.5]).unwrap();
19418        p
19419    }
19420
19421    #[test]
19422    fn mimo_embedded_prompt_uses_rows_and_never_reuses_token_only_kv() {
19423        let mut p = mimo_test_pipeline();
19424        p.speculative = false;
19425        p.ignore_eos = true;
19426        p.sampler_config.temperature = 0.0;
19427        p.sampler_config.repetition_penalty = 1.0;
19428        let a = vec![3, 5, 7, 9, 11, 13];
19429        let b = vec![4, 8, 12, 16, 20, 24];
19430        let rows: Vec<_> = b.iter().flat_map(|&id| p.embed_id(id)).collect();
19431        let expected = p.generate_from_ids(&b, 8, None, None).unwrap().token_ids;
19432        // Same placeholder IDs as an earlier request are not a cache key
19433        // for different media. The actual rows, not a re-embedding of a,
19434        // must determine the continuation.
19435        let actual = p.generate_from_embeds(&a, &rows, 8, None, None).unwrap().token_ids;
19436        assert_eq!(actual, expected);
19437        assert!(p.kv_history.is_empty());
19438        let mut extended = a.clone();
19439        extended.push(17);
19440        let after_media = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19441        p.reset_session();
19442        let fresh = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19443        assert_eq!(after_media, fresh);
19444        assert!(p.generate_from_embeds(&a, &rows[..rows.len()-1], 1, None, None).is_err());
19445        // Force a real token-prefix reuse opportunity into the media call.
19446        // Those labels are unchanged, but their embeddings now describe a
19447        // different source sequence and every KV row must be rebuilt.
19448        p.reset_session();
19449        p.generate_from_ids(&a, 1, None, None).unwrap();
19450        let mut media_ids = p.kv_history.clone();
19451        assert!(!media_ids.is_empty());
19452        media_ids.extend_from_slice(&[19, 21, 23]);
19453        let source_ids: Vec<_> = (0..media_ids.len()).map(|i| b[i % b.len()]).collect();
19454        let source_rows: Vec<_> = source_ids.iter().flat_map(|&id| p.embed_id(id)).collect();
19455        let mut oracle = mimo_test_pipeline();
19456        oracle.speculative = false;
19457        oracle.ignore_eos = true;
19458        oracle.sampler_config.temperature = 0.0;
19459        oracle.sampler_config.repetition_penalty = 1.0;
19460        let expected = oracle.generate_from_ids(&source_ids, 8, None, None).unwrap().token_ids;
19461        assert_eq!(p.generate_from_embeds(&media_ids, &source_rows, 8, None, None).unwrap().token_ids, expected);
19462        assert!(p.kv_history.is_empty());
19463        let mut bad = rows;
19464        bad[0] = f32::NAN;
19465        assert!(p.generate_from_embeds(&a, &bad, 1, None, None).is_err());
19466    }
19467
19468    fn f32_bits(v: &[f32]) -> Vec<u32> {
19469        v.iter().map(|x| x.to_bits()).collect()
19470    }
19471
19472    /// M3 acceptance: on the MiMo-shaped stack the decode walk (one
19473    /// position at a time through `forward_layers`) and the batched
19474    /// prefill (`prefill_batch_span`, whole prompt and split in two
19475    /// chunks) give bit-identical logits at all 12 positions — per-layer
19476    /// KV heads, narrow V, sinks, the window and the biased sigmoid MoE all
19477    /// agree across the two walks.
19478    #[test]
19479    fn mimo_shaped_decode_matches_prefill_batch_bitwise() {
19480        let mut p = mimo_test_pipeline();
19481        let kv: Vec<usize> = p.kv_cache.layers.iter().map(|l| l.num_kv_heads).collect();
19482        assert_eq!(kv, vec![1, 2, 2, 1]);
19483        assert!(p.kv_cache.layers[0].sinks.is_none() && p.kv_cache.layers[3].sinks.is_none());
19484        assert!(p.kv_cache.layers[1].sinks.is_some() && p.kv_cache.layers[2].sinks.is_some());
19485        let ids: Vec<u32> = (0..12u32).map(|i| (i * 7 + 3) % 64).collect();
19486        let hs = p.hidden_size;
19487        let mut decode = Vec::new();
19488        for (pos, &id) in ids.iter().enumerate() {
19489            let e = p.embed_single(id);
19490            let h = p.forward_layers(&e, pos, None);
19491            decode.push(p.logits_from_hidden(&h));
19492        }
19493        for l in &p.kv_cache.layers {
19494            assert_eq!(l.seq_len, 12);
19495            // V rows are padded to head_dim inside the cache.
19496            assert_eq!(l.head_values(0).len(), 12 * 8);
19497        }
19498        assert!(decode.iter().flatten().all(|v| v.is_finite()));
19499
19500        p.clear_sequence_state();
19501        let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19502        for pos in 0..ids.len() {
19503            let lg = p.logits_from_hidden(&hb[pos * hs..(pos + 1) * hs]);
19504            assert_eq!(
19505                f32_bits(&decode[pos]),
19506                f32_bits(&lg),
19507                "whole prompt, pos {pos}"
19508            );
19509        }
19510
19511        p.clear_sequence_state();
19512        let a = p.prefill_batch_span(PrefillIn::Ids(&ids[..5]), 0, None, 0, p.num_layers);
19513        let b = p.prefill_batch_span(PrefillIn::Ids(&ids[5..]), 5, None, 0, p.num_layers);
19514        for pos in 0..ids.len() {
19515            let row = if pos < 5 {
19516                &a[pos * hs..(pos + 1) * hs]
19517            } else {
19518                &b[(pos - 5) * hs..(pos - 4) * hs]
19519            };
19520            let lg = p.logits_from_hidden(row);
19521            assert_eq!(
19522                f32_bits(&decode[pos]),
19523                f32_bits(&lg),
19524                "two chunks, pos {pos}"
19525            );
19526        }
19527
19528        // The fixture is not degenerate: the sinks and the window each
19529        // change the answer.
19530        let last = |p: &mut Pipeline| {
19531            p.clear_sequence_state();
19532            let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19533            p.logits_from_hidden(&hb[11 * hs..12 * hs])
19534        };
19535        let base = last(&mut p);
19536        let mut no_sinks = mimo_test_pipeline();
19537        for l in &mut no_sinks.kv_cache.layers {
19538            l.sinks = None;
19539        }
19540        assert_ne!(
19541            f32_bits(&last(&mut no_sinks)),
19542            f32_bits(&base),
19543            "sinks are live"
19544        );
19545        let mut wide = mimo_test_pipeline();
19546        wide.swa = Some((64, usize::MAX));
19547        assert_ne!(
19548            f32_bits(&last(&mut wide)),
19549            f32_bits(&base),
19550            "window is live"
19551        );
19552
19553        // Generation runs end to end on the same stack.
19554        p.clear_sequence_state();
19555        p.ignore_eos = true;
19556        let r = p.generate_from_ids(&ids, 4, None, None).unwrap();
19557        assert_eq!(r.token_ids.len(), 4);
19558    }
19559
19560    /// A synthetic MiMo draft stack of `n` layers for `mimo_test_pipeline`
19561    /// (the SWA geometry of its sliding layers: 2 KV heads, head 8 / V 4).
19562    fn mimo_test_mtp(n: usize, gain: f32) -> mimo_mtp::MimoMtp {
19563        let (hs, inter, nh, hd, vd, nkv) = (16usize, 24usize, 4usize, 8usize, 4usize, 2usize);
19564        let synth = |len: usize, salt: usize| -> Vec<f32> {
19565            (0..len)
19566                .map(|i| (((i * 37 + salt * 13 + 3) % 89) as f32 / 89.0 - 0.5) * 0.6 * gain)
19567                .collect()
19568        };
19569        let qt = |rows: usize, cols: usize, salt: usize| {
19570            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19571        };
19572        let layers = (0..n)
19573            .map(|k| {
19574                let s = 500 + k * 40;
19575                let mut kv = crate::kv_cache::LayerKvCache::new(nkv, hd);
19576                kv.sinks = Some(vec![0.3, -0.7, 1.1, 0.0]);
19577                MtpModule {
19578                    enorm: vec![1.0; hs],
19579                    hnorm: vec![1.0; hs],
19580                    eh_proj: qt(hs, 2 * hs, s),
19581                    layer: LayerWeights {
19582                        input_norm: vec![1.0; hs],
19583                        post_norm: vec![1.0; hs],
19584                        attn_out_norm: None,
19585                        ffn_out_norm: None,
19586                        layer_scale: None,
19587                        attn: AttnKind::Full {
19588                            wq: qt(nh * hd, hs, s + 1),
19589                            wk: qt(nkv * hd, hs, s + 2),
19590                            wv: qt(nkv * vd, hs, s + 3),
19591                            wo: qt(hs, nh * vd, s + 4),
19592                            q_norm: None,
19593                            k_norm: None,
19594                            output_gate: false,
19595                            softplus_gate: None,
19596                            bias: None,
19597                        },
19598                        ffn: FfnKind::Dense(DenseFfn {
19599                            gate_proj: qt(inter, hs, s + 5),
19600                            up_proj: qt(inter, hs, s + 6),
19601                            down_proj: qt(hs, inter, s + 7),
19602                            act: Act::Silu,
19603                            down_t: None,
19604                            segs: Vec::new(),
19605                        }),
19606                    },
19607                    final_norm: vec![1.0; hs],
19608                    kv,
19609                }
19610            })
19611            .collect();
19612        mimo_mtp::MimoMtp::from_layers(layers)
19613    }
19614
19615    fn mimo_greedy(p: &mut Pipeline, ids: &[u32], n: usize, spec: bool) -> GenerateResult {
19616        p.clear_sequence_state();
19617        p.speculative = spec;
19618        p.ignore_eos = true;
19619        p.sampler_config.temperature = 0.0;
19620        p.generate_from_ids(ids, n, None, None).unwrap()
19621    }
19622
19623    /// The draft stack's incremental rounds (a few rows per layer, last
19624    /// round's provisional rows dropped) give exactly the teacher-forced
19625    /// table of one causal pass per layer over the whole sequence — the
19626    /// table `tools/mimo_ref.py mtp` computes for variant A: layer k, row
19627    /// j reads (x[j+k+1], norm(h_j)) at RoPE position j.
19628    #[test]
19629    fn mimo_mtp_incremental_rounds_equal_the_teacher_forced_table() {
19630        // Both readings of the backbone hidden: pre-final-norm (default)
19631        // and post-final-norm (`CMF_MIMO_MTP_HIDDEN=post`).
19632        for post in [false, true] {
19633            let mut p = mimo_test_pipeline();
19634            // A non-trivial final norm, so the two readings differ.
19635            p.weights.final_norm = (0..p.hidden_size).map(|i| 0.5 + 0.1 * i as f32).collect();
19636            let mut st0 = mimo_test_mtp(3, 1.0);
19637            st0.post_norm_hidden = post;
19638            p.mimo_mtp = Some(st0);
19639            let ids: Vec<u32> = (0..14u32).map(|i| (i * 11 + 5) % 64).collect();
19640            let hs = p.hidden_size;
19641            let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19642            p.mimo_note_rows(&hb, 0);
19643            let mut st = p.mimo_mtp.take().unwrap();
19644            // Incremental: one round per t through the decode path (later
19645            // tokens from `ids`, the probe's teacher forcing).
19646            let k = 3;
19647            let mut inc = Vec::new();
19648            for t in 0..ids.len() - k - 1 {
19649                inc.push(p.mimo_mtp_draft(&mut st, t, &ids, k));
19650            }
19651            // Reference: per layer, ONE batched causal pass over all rows
19652            // with fresh caches.
19653            let s = ids.len();
19654            let mut reference = vec![vec![0u32; k]; s - k - 1];
19655            let mut fresh = mimo_test_mtp(3, 1.0);
19656            for (layer, m) in fresh.layers.iter_mut().enumerate() {
19657                let n = s - layer - 1;
19658                let mut cats = vec![0.0f32; n * 2 * hs];
19659                for j in 0..n {
19660                    let e = p.embed_single(ids[j + layer + 1]);
19661                    let raw = &hb[j * hs..(j + 1) * hs];
19662                    let g = if post {
19663                        inference::rms_norm(raw, &p.weights.final_norm, p.rms_eps, p.norm_style)
19664                    } else {
19665                        raw.to_vec()
19666                    };
19667                    let (ce, ch) = cats[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
19668                    inference::rms_norm_into(&e, &m.enorm, p.rms_eps, p.norm_style, ce);
19669                    inference::rms_norm_into(&g, &m.hnorm, p.rms_eps, p.norm_style, ch);
19670                }
19671                let mut x = vec![0.0f32; n * hs];
19672                m.eh_proj.matmat(&cats, n, &mut x, None);
19673                p.mimo_mtp_block(m, &mut x, n, 0);
19674                for (t, row) in reference.iter_mut().enumerate() {
19675                    let y = inference::rms_norm(
19676                        &x[t * hs..(t + 1) * hs],
19677                        &m.final_norm,
19678                        p.rms_eps,
19679                        p.norm_style,
19680                    );
19681                    row[layer] = sampler::argmax(&p.lm_head_forward(&y));
19682                }
19683            }
19684            assert_eq!(inc, reference, "post_norm_hidden = {post}");
19685            // Not a degenerate table: the drafts vary.
19686            let distinct: std::collections::HashSet<u32> =
19687                inc.iter().flatten().copied().collect();
19688            assert!(distinct.len() > 3, "{inc:?}");
19689            // Each layer's cache ends holding rows up to the last round start.
19690            let last_t = ids.len() - k - 2;
19691            for m in &st.layers {
19692                assert_eq!(m.kv.seq_len, last_t + 1);
19693            }
19694        }
19695    }
19696
19697    /// Greedy with the MiMo draft stack is the plain greedy stream, token
19698    /// for token — with the real draft layers (low acceptance) and with a
19699    /// drafter that is right most of the time (exercises accepted prefixes
19700    /// of every length, the KV truncation of the rejected rows and the
19701    /// logits hand-off to the loop top), under the default repetition
19702    /// penalty.
19703    #[test]
19704    fn mimo_speculative_greedy_equals_plain_greedy() {
19705        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19706        let ids: Vec<u32> = (0..9u32).map(|i| (i * 7 + 3) % 64).collect();
19707        let n = 24;
19708        let mut p = mimo_test_pipeline();
19709        let plain = mimo_greedy(&mut p, &ids, n, false);
19710        assert_eq!(plain.mtp_drafted, 0);
19711        assert_eq!(plain.token_ids.len(), n);
19712        let plain_kv = p.kv_cache.layers[0].seq_len;
19713
19714        // Real draft layers.
19715        p.mimo_mtp = Some(mimo_test_mtp(3, 1.0));
19716        let spec = mimo_greedy(&mut p, &ids, n, true);
19717        assert!(spec.mtp_drafted > 0, "the round must draft");
19718        assert_eq!(spec.token_ids, plain.token_ids);
19719        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19720
19721        // A drafter reading the true continuation with every fifth token
19722        // wrong: accepted prefixes of 0..=3 all occur.
19723        let mut truth: Vec<u32> = ids.clone();
19724        truth.extend(&plain.token_ids);
19725        let mut noisy = truth.clone();
19726        for (i, t) in noisy.iter_mut().enumerate() {
19727            if i % 5 == 0 {
19728                *t = (*t + 1) % 64;
19729            }
19730        }
19731        let mut st = mimo_test_mtp(3, 1.0);
19732        st.draft_override = Some(noisy);
19733        p.mimo_mtp = Some(st);
19734        let spec = mimo_greedy(&mut p, &ids, n, true);
19735        assert_eq!(spec.token_ids, plain.token_ids);
19736        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19737        let stats = p.mimo_mtp.as_ref().unwrap().stats.clone();
19738        assert_eq!(stats.accepted as usize, spec.mtp_accepted);
19739        assert!(spec.mtp_accepted > 0 && spec.mtp_accepted < spec.mtp_drafted);
19740        assert!(
19741            stats.accept_hist.iter().filter(|&&c| c > 0).count() >= 3,
19742            "{:?}",
19743            stats.accept_hist
19744        );
19745        assert!(stats.tokens_per_round() > 1.5, "{}", stats.line());
19746
19747        // A perfect drafter: every draft accepted, rounds of K+1 tokens,
19748        // and the budget is never overrun.
19749        let mut st = mimo_test_mtp(3, 1.0);
19750        st.draft_override = Some(truth);
19751        p.mimo_mtp = Some(st);
19752        let spec = mimo_greedy(&mut p, &ids, n, true);
19753        assert_eq!(spec.token_ids, plain.token_ids);
19754        assert_eq!(spec.mtp_accepted, spec.mtp_drafted);
19755        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19756
19757        // CMF_MTP=0 path: the stack is attached but idle.
19758        let off = mimo_greedy(&mut p, &ids, n, false);
19759        assert_eq!(off.token_ids, plain.token_ids);
19760        assert_eq!(off.mtp_drafted, 0);
19761    }
19762
19763    /// The wgpu graphs carry MiMo-V2's attention per layer (KV heads,
19764    /// narrow V, sinks, windows, two RoPE tables): no attention-level
19765    /// decline for it any more, and the geometry each layer hands the
19766    /// graph is exactly what the CPU attention reads for that layer. The
19767    /// descriptive reasons stay (the Metal graphs and the q1 dropin still
19768    /// decline on them), and what the per-layer geometry cannot express
19769    /// keeps a named wgpu decline.
19770
19771    #[test]
19772    fn mimo_shaped_model_rides_the_wgpu_graph_geometry() {
19773        let p = mimo_test_pipeline();
19774        assert_eq!(
19775            p.graph_attn_decline_reason(),
19776            Some("per-layer KV head counts")
19777        );
19778        assert_eq!(p.wgpu_graph_attn_decline(), None);
19779        let g0 = p.graph_attn_geom(0).expect("full layer geometry");
19780        assert_eq!(
19781            (g0.nkv, g0.dv, g0.rd, g0.window, g0.sink.is_some()),
19782            (1, 4, 4, None, false)
19783        );
19784        assert_eq!(g0.invf, p.inv_freq.as_slice());
19785        let g1 = p.graph_attn_geom(1).expect("sliding layer geometry");
19786        assert_eq!((g1.nkv, g1.dv, g1.rd, g1.window), (2, 4, 4, Some(3)));
19787        assert_eq!(g1.sink, Some(&[0.5f32, -1.0, 1.5, 0.0][..]));
19788        assert_eq!(g1.invf, p.inv_freq_local.as_ref().unwrap().as_slice());
19789        assert_ne!(g0.invf, g1.invf, "two RoPE tables");
19790        let g3 = p.graph_attn_geom(3).expect("full layer geometry");
19791        assert_eq!((g3.nkv, g3.window, g3.sink.is_some()), (1, None, false));
19792
19793        // No wgpu device in this process: the builders run and decline on
19794        // the (f32, unmapped) experts — never with an attention line.
19795        let emb = p.embed_single(3);
19796        let mut lg = Vec::new();
19797        assert!(
19798            p.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, p.num_layers)
19799                .is_none()
19800        );
19801        let mut hid = emb.clone();
19802        assert_eq!(
19803            p.try_batch_graph_wgpu(&mut hid, &[0], 1, None),
19804            crate::gpu::BatchGraphOutcome::Declined
19805        );
19806        assert_eq!(hid, emb, "a declined batch graph leaves the rows untouched");
19807        assert!(p.try_multi_burst(3, 0, 4).is_none());
19808        assert!(
19809            p.graph_declines().is_empty(),
19810            "no attention decline logged: {:?}",
19811            p.graph_declines()
19812        );
19813        // (No assertion on graph_prefill_preferred: with no attention
19814        // decline it follows the device — a test process that brought a
19815        // wgpu adapter up routes this resident MoE through the graph.)
19816
19817        let plain = || create_test_pipeline(8, 16, 2, 1, 4, 2, 32);
19818        assert_eq!(plain().graph_attn_decline_reason(), None);
19819        assert_eq!(plain().wgpu_graph_attn_decline(), None);
19820        assert!(
19821            plain().graph_attn_geom(0).is_none(),
19822            "uniform models keep the historical arms"
19823        );
19824        let mut q = plain();
19825        q.set_layer_sinks(1, vec![0.25, -0.25]).unwrap();
19826        assert_eq!(
19827            q.graph_attn_decline_reason(),
19828            Some("learned attention sinks")
19829        );
19830        assert_eq!(
19831            q.graph_attn_geom(1).unwrap().sink,
19832            Some(&[0.25f32, -0.25][..])
19833        );
19834        let mut q = plain();
19835        q.set_attn_geometry(None, Some(2)).unwrap();
19836        assert_eq!(
19837            q.graph_attn_decline_reason(),
19838            Some("V heads narrower than Q/K heads")
19839        );
19840        assert_eq!(q.graph_attn_geom(0).unwrap().dv, 2);
19841        let mut q = plain();
19842        q.sliding_layers = Some(vec![true, false]);
19843        q.swa = Some((4, usize::MAX));
19844        assert_eq!(q.graph_attn_decline_reason(), Some("sliding-window layers"));
19845        assert_eq!(q.graph_attn_geom(0).unwrap().window, Some(4));
19846        assert_eq!(q.graph_attn_geom(1).unwrap().window, None);
19847
19848        // Outside the per-layer geometry: a named wgpu decline, logged
19849        // once per site.
19850        let mut q = mimo_test_pipeline();
19851        q.rope_scale = 2.0;
19852        assert_eq!(
19853            q.wgpu_graph_attn_decline(),
19854            Some("scaled RoPE positions with per-layer geometry")
19855        );
19856        let emb = q.embed_single(3);
19857        assert!(
19858            q.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, q.num_layers)
19859                .is_none()
19860        );
19861        let _ = q.try_token_graph_wgpu_steps(&emb, 1, &mut lg, 1, None, None, 0, q.num_layers);
19862        let lines = q.graph_declines();
19863        assert_eq!(
19864            lines
19865                .iter()
19866                .filter(|l| l.starts_with("wgpu token graph") && l.contains("scaled RoPE"))
19867                .count(),
19868            1,
19869            "{lines:?}"
19870        );
19871    }
19872
19873    #[test]
19874    fn mimo_verify_rewind_preserves_lagging_host_caches() {
19875        let mut p = mimo_test_pipeline();
19876        for (li, layer) in p.kv_cache.layers.iter_mut().enumerate() {
19877            let row = vec![0.0; layer.num_kv_heads * layer.head_dim];
19878            for _ in 0..if li == 0 { 2 } else { 12 } {
19879                layer.append(&row, &row, &[]);
19880            }
19881        }
19882        p.mimo_verify_rewind(9).unwrap();
19883        assert_eq!(p.kv_cache.layers[0].seq_len, 2);
19884        for layer in &p.kv_cache.layers[1..] {
19885            assert_eq!(layer.seq_len, 9);
19886        }
19887    }
19888
19889    /// CMF_LAYER_DUMP: the decode walk and the batched prefill both write
19890    /// every (position, layer) hidden, the two sets agree byte for byte,
19891    /// and the last layer's file is the stack output.
19892    #[test]
19893    fn layer_dump_covers_every_position_and_layer_on_both_walks() {
19894        let dir = std::env::temp_dir().join(format!("cmf-layer-dump-{}", std::process::id()));
19895        let _ = std::fs::remove_dir_all(&dir);
19896        let mut p = mimo_test_pipeline();
19897        let hs = p.hidden_size;
19898        let ids = [5u32, 9, 11, 2, 40];
19899        p.layer_dump = Some(dir.join("decode"));
19900        for (pos, &id) in ids.iter().enumerate() {
19901            let e = p.embed_single(id);
19902            let _ = p.forward_layers(&e, pos, None);
19903        }
19904        p.clear_sequence_state();
19905        p.layer_dump = Some(dir.join("prefill"));
19906        let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19907        for pos in 0..ids.len() {
19908            for li in 0..p.num_layers {
19909                let name = format!("p{pos:06}_l{li:02}.f32");
19910                let a = std::fs::read(dir.join("decode").join(&name)).unwrap();
19911                let b = std::fs::read(dir.join("prefill").join(&name)).unwrap();
19912                assert_eq!(a.len(), hs * 4, "{name}");
19913                assert_eq!(a, b, "{name}");
19914            }
19915        }
19916        let last = std::fs::read(dir.join("prefill").join("p000004_l03.f32")).unwrap();
19917        let vals: Vec<f32> = last
19918            .chunks(4)
19919            .map(|c| f32::from_le_bytes(c.try_into().unwrap()))
19920            .collect();
19921        assert_eq!(f32_bits(&vals), f32_bits(&hb[4 * hs..5 * hs]));
19922        let _ = std::fs::remove_dir_all(&dir);
19923    }
19924
19925    #[test]
19926    fn attn_geometry_and_sinks_are_validated() {
19927        let mut p = create_test_pipeline(8, 16, 4, 2, 4, 2, 32);
19928        assert!(
19929            p.set_attn_geometry(Some(vec![2]), None).is_err(),
19930            "one entry per layer"
19931        );
19932        assert!(
19933            p.set_attn_geometry(Some(vec![2, 3]), None).is_err(),
19934            "3 does not divide 4"
19935        );
19936        assert!(p.set_attn_geometry(Some(vec![2, 0]), None).is_err());
19937        assert!(p.set_attn_geometry(None, Some(0)).is_err());
19938        assert!(
19939            p.set_attn_geometry(None, Some(5)).is_err(),
19940            "V wider than the head"
19941        );
19942        p.set_attn_geometry(None, Some(4)).unwrap();
19943        assert_eq!(
19944            p.v_head_dim, None,
19945            "v_head_dim == head_dim is the uniform case"
19946        );
19947        p.set_layer_sinks(1, vec![0.1; 4]).unwrap();
19948        p.set_attn_geometry(Some(vec![1, 4]), None).unwrap();
19949        assert_eq!(p.kv_cache.layers[0].num_kv_heads, 1);
19950        assert_eq!(p.kv_cache.layers[1].num_kv_heads, 4);
19951        assert!(
19952            p.kv_cache.layers[1].sinks.is_some(),
19953            "a reshape keeps the layer's sinks"
19954        );
19955        assert_eq!(p.layer_geom(1).0, 4);
19956        assert!(
19957            p.set_layer_sinks(0, vec![0.0; 3]).is_err(),
19958            "one sink per Q head"
19959        );
19960        assert!(p.set_layer_sinks(7, vec![0.0; 4]).is_err());
19961        assert!(p.set_layer_sinks(0, vec![f32::NAN, 0.0, 0.0, 0.0]).is_err());
19962    }
19963
19964    /// The O(1) Nyström state replaces a plain full-context softmax; it
19965    /// must never be armed on a sliding, sink or narrow-V layer.
19966    #[test]
19967    fn o1_is_never_armed_on_sink_window_or_narrow_v_layers() {
19968        let cfg = || {
19969            Some(crate::nystrom::O1Cfg {
19970                layers: crate::nystrom::O1Layers::All,
19971                m: 4,
19972                w: 8,
19973                sink: 2,
19974                rect: crate::nystrom::O1Rect::Aggregate,
19975            })
19976        };
19977        let mut p = mimo_test_pipeline();
19978        p.set_o1(cfg());
19979        assert!(!p.o1_active(), "every MiMo-shaped layer is ineligible");
19980        let mut q = create_test_pipeline(8, 16, 2, 1, 4, 3, 64);
19981        q.set_layer_sinks(1, vec![0.0, 0.0]).unwrap();
19982        q.sliding_layers = Some(vec![false, false, true]);
19983        q.swa = Some((4, usize::MAX));
19984        q.set_o1(cfg());
19985        assert_eq!(q.o1_flags, vec![true, false, false]);
19986    }
19987
19988    #[test]
19989    fn o1_batch_transition_publishes_one_epoch_before_serial_handoff() {
19990        const B: usize = 19;
19991        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19992        p.set_o1(Some(crate::nystrom::O1Cfg {
19993            layers: crate::nystrom::O1Layers::All,
19994            m: 4,
19995            w: 8,
19996            sink: 2,
19997            rect: crate::nystrom::O1Rect::Aggregate,
19998        }));
19999        p.o1_begin_with_prefix(Some(B));
20000        let ids: Vec<u32> = (0..B as u32).collect();
20001        let _ = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
20002
20003        assert_eq!(p.o1_epoch, 1, "all layers publish one completed transition");
20004        assert!(p.kv_cache.layers.iter().all(|l| l.o1_sealed()));
20005        let next = p.embed_single(B as u32);
20006        let _ = p.forward_layers(&next, B, None);
20007        assert_eq!(p.o1_epoch, 1, "sealed handoff must not republish the epoch");
20008    }
20009
20010    #[test]
20011    fn o1_pair_transition_commits_scratch_before_epoch_publication() {
20012        const B: usize = 19;
20013        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
20014        // Keep a real recurrent layer ahead of the Full O(1) layer so the
20015        // pair test observes the GDN lane-2 scratch swap at the same
20016        // boundary, rather than only exercising an artificial scratch vec.
20017        let gdn_cfg = crate::linear_core::GdnCfg {
20018            num_v_heads: 2,
20019            num_k_heads: 1,
20020            key_head_dim: 2,
20021            value_head_dim: 4,
20022            conv_kernel: 3,
20023            hidden_size: 8,
20024            rms_eps: 1e-6,
20025            output_gate_sigmoid: false,
20026        };
20027        let synth = |n: usize, salt: usize| -> Vec<f32> {
20028            (0..n)
20029                .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
20030                .collect()
20031        };
20032        let qt = |rows: usize, cols: usize, salt: usize| {
20033            crate::qtensor::QTensor::from_f32(synth(rows * cols, salt), rows, cols)
20034        };
20035        let c_dim = gdn_cfg.conv_dim();
20036        let vd = gdn_cfg.num_v_heads * gdn_cfg.value_head_dim;
20037        p.weights.layers[0].attn = AttnKind::LinearGdn(crate::linear_core::GdnWeights {
20038            in_proj_qkv: qt(c_dim, 8, 1),
20039            in_proj_z: qt(vd, 8, 2),
20040            in_proj_a: qt(gdn_cfg.num_v_heads, 8, 3),
20041            in_proj_b: qt(gdn_cfg.num_v_heads, 8, 4),
20042            conv1d: synth(c_dim * gdn_cfg.conv_kernel, 5),
20043            a_log: vec![0.2, 0.5],
20044            dt_bias: synth(gdn_cfg.num_v_heads, 6),
20045            norm: vec![1.0; gdn_cfg.value_head_dim],
20046            out_proj: qt(8, vd, 7),
20047        });
20048        p.gdn_cfg = Some(gdn_cfg);
20049        p.set_o1(Some(crate::nystrom::O1Cfg {
20050            layers: crate::nystrom::O1Layers::All,
20051            m: 4,
20052            w: 8,
20053            sink: 2,
20054            rect: crate::nystrom::O1Rect::Aggregate,
20055        }));
20056        p.o1_begin_with_prefix(Some(B));
20057        for pos in 0..B - 2 {
20058            let emb = p.embed_single(pos as u32);
20059            let _ = p.forward_layers(&emb, pos, None);
20060        }
20061        let lane1_state = p.kv_cache.layers[0].linear_state.clone();
20062
20063        let e1 = p.embed_single((B - 2) as u32);
20064        let e2 = p.embed_single((B - 1) as u32);
20065        let _ = p.forward_pair(&e1, &e2, B - 2);
20066
20067        assert_eq!(p.o1_epoch, 1, "pair crossing B publishes one epoch");
20068        assert!(
20069            p.kv_cache
20070                .layers
20071                .iter()
20072                .enumerate()
20073                .all(|(li, l)| !p.o1_flags[li] || l.o1_sealed())
20074        );
20075        assert!(!p.kv_cache.layers[0].linear_state.is_empty());
20076        assert_ne!(
20077            p.kv_cache.layers[0].linear_state, lane1_state,
20078            "real pair must commit GDN lane 2 before returning"
20079        );
20080        assert!(p.kv_cache.layers[0].linear_scratch.is_empty());
20081        let next = p.embed_single(B as u32);
20082        let _ = p.forward_layers(&next, B, None);
20083        assert_eq!(p.o1_epoch, 1, "serial continuation must reuse the epoch");
20084    }
20085
20086    #[test]
20087    fn o1_error_observation_stays_terminal_until_reset() {
20088        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20089        p.set_o1(Some(crate::nystrom::O1Cfg {
20090            layers: crate::nystrom::O1Layers::All,
20091            m: 4,
20092            w: 8,
20093            sink: 2,
20094            rect: crate::nystrom::O1Rect::Aggregate,
20095        }));
20096        p.o1_begin();
20097        p.kv_cache.layers[0].o1_abort("synthetic transition failure".into());
20098
20099        assert!(p.o1_seal_checked().is_err());
20100        assert!(
20101            p.o1_seal_checked().is_err(),
20102            "retry must see the sticky error"
20103        );
20104        let k = vec![0.2f32; 4];
20105        let v = vec![0.3f32; 4];
20106        p.kv_cache.layers[0].append(&k, &v, &[]);
20107        assert_eq!(p.kv_cache.layers[0].seq_len, 0);
20108
20109        p.reset_session();
20110        p.o1_begin();
20111        p.kv_cache.layers[0].append(&k, &v, &[]);
20112        assert_eq!(p.kv_cache.layers[0].seq_len, 1);
20113    }
20114
20115    #[test]
20116    fn nll_graph_failure_is_terminal_and_request_is_reusable() {
20117        let ids = vec![1u32, 2, 3, 4, 5, 6];
20118        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20119        p.graph_logits = Some(vec![123.0]);
20120        p.graph_want_logits = true;
20121        p.graph_failed
20122            .store(true, std::sync::atomic::Ordering::Relaxed);
20123        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
20124        let err = p.nll_ids_from(&ids, 0).expect_err("prior graph failure");
20125        assert!(err.contains("before NLL"));
20126        assert!(p.graph_logits.is_none());
20127        assert!(!p.graph_want_logits);
20128        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20129        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20130
20131        let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20132        let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
20133        let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
20134        assert_eq!(actual.1, expected.1);
20135        assert!((actual.0 - expected.0).abs() < 1e-9);
20136    }
20137
20138    #[test]
20139    fn nll_forward_failure_discards_partial_score_and_clears_sidechannels() {
20140        let ids = vec![1u32, 2, 3, 4, 5, 6];
20141        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20142        p.nll_test_fail_at = Some(1);
20143        let err = p
20144            .nll_ids_from(&ids, 0)
20145            .expect_err("one-shot forward failure");
20146        assert!(err.contains("forward") || err.contains("score row"));
20147        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20148        assert!(!p.graph_want_logits);
20149        assert!(p.graph_logits.is_none());
20150        assert!(p.kv_history.is_empty());
20151
20152        let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20153        let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
20154        let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
20155        assert_eq!(actual.1, expected.1);
20156        assert!((actual.0 - expected.0).abs() < 1e-9);
20157    }
20158
20159    #[test]
20160    fn nll_serial_failure_before_first_row_is_reported() {
20161        let ids = vec![1u32, 2, 3, 4];
20162        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20163        p.nll_test_force_serial = true;
20164        p.nll_test_fail_at = Some(0);
20165        let err = p.nll_ids_from(&ids, 0).expect_err("serial forward failure");
20166        assert!(err.contains("serial forward"));
20167        assert!(p.kv_history.is_empty());
20168        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20169        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20170    }
20171
20172    #[test]
20173    fn ffn_probe_failure_discards_recorder_and_state() {
20174        let ids = vec![1u32, 2, 3, 4];
20175        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20176        p.nll_test_fail_at = Some(0);
20177        let err = p
20178            .probe_ffn_mass_batch(&ids)
20179            .expect_err("probe forward failure");
20180        assert!(err.contains("NLL"));
20181        assert!(FFN_PROBE.with(|probe| probe.borrow().is_none()));
20182        assert!(p.kv_history.is_empty());
20183        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20184    }
20185
20186    #[test]
20187    fn nll_test_controls_are_pipeline_scoped() {
20188        let ids = vec![1u32, 2, 3, 4];
20189        let mut failing = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20190        let mut unaffected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20191        failing.nll_test_force_serial = true;
20192        failing.nll_test_fail_at = Some(0);
20193
20194        assert!(!failing.can_prefill_batched());
20195        assert!(unaffected.can_prefill_batched());
20196        let expected = unaffected
20197            .nll_ids_from(&ids, 0)
20198            .expect("unaffected pipeline remains usable");
20199        let err = failing
20200            .nll_ids_from(&ids, 0)
20201            .expect_err("failure injection belongs to failing pipeline");
20202        assert!(err.contains("serial forward"));
20203        assert!(failing.nll_test_fail_at.is_none());
20204        assert!(unaffected.can_prefill_batched());
20205        let actual = unaffected
20206            .nll_ids_from(&ids, 0)
20207            .expect("unaffected pipeline remains reusable");
20208        assert_eq!(actual.1, expected.1);
20209        assert!((actual.0 - expected.0).abs() < 1e-9);
20210    }
20211
20212    #[test]
20213    fn forward_ids_failure_channel_is_terminal_and_reusable() {
20214        let ids = vec![1u32, 2, 3, 4, 5, 6];
20215        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20216        p.graph_logits = Some(vec![123.0]);
20217        p.graph_want_logits = true;
20218        p.graph_failed
20219            .store(true, std::sync::atomic::Ordering::Relaxed);
20220        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
20221
20222        let err = p
20223            .forward_ids(&ids, None)
20224            .expect_err("a failed forward must not become a valid head result");
20225        assert!(err.contains("forward_ids setup"));
20226        assert!(p.graph_logits.is_none());
20227        assert!(!p.graph_want_logits);
20228        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20229        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20230        assert_eq!(p.kv_cache.seq_len(), 0);
20231
20232        let expected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64)
20233            .forward_ids(&ids, None)
20234            .expect("fresh forward_ids");
20235        let actual = p
20236            .forward_ids(&ids, None)
20237            .expect("pipeline remains reusable after a failed forward");
20238        assert_eq!(actual.len(), expected.len());
20239        assert!(
20240            actual
20241                .iter()
20242                .zip(expected)
20243                .all(|(a, b)| (a - b).abs() < 1e-9)
20244        );
20245        assert_eq!(p.kv_cache.seq_len(), ids.len());
20246    }
20247
20248    #[test]
20249    fn sigmoid_router_floor_is_explicit_per_architecture() {
20250        // GLM-5's noaux_tc reference uses +1e-20 while the generic
20251        // LFM2-compatible path uses +1e-6.  At low (but representable)
20252        // sigmoid scores, silently sharing the latter changes expert weights
20253        // by orders of magnitude and can make a routed layer look coherent
20254        // while discarding its expert contribution.
20255        let zero = || DenseFfn {
20256            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20257            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20258            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20259            act: Act::Silu,
20260            down_t: None,
20261            segs: Vec::new(),
20262        };
20263        let m = MoeFfn {
20264            router: QTensor::from_f32(vec![0.0; 4], 2, 2),
20265            experts: vec![zero(), zero()],
20266            top_k: 1,
20267            norm_topk_prob: true,
20268            router_sigmoid: true,
20269            expert_bias: None,
20270            routed_scaling: 2.5,
20271            route_tau: None,
20272            shared: None,
20273            stats: std::cell::RefCell::new(Vec::new()),
20274            act_sq: std::cell::RefCell::new(Vec::new()),
20275            act_rows: std::cell::RefCell::new(Vec::new()),
20276            mask: None,
20277            per_expert_scale: None,
20278            router_input_norm: false,
20279            resonance: None,
20280            grown: Vec::new(),
20281        };
20282        let logits = [-20.0f32, -20.0];
20283        let (_, p, glm_wsum) = moe_route_with_eps(&logits, &m, None, 1e-20);
20284        let (_, _, generic_wsum) = moe_route(&logits, &m, None);
20285        let expected = (p[0] + 1e-20) / m.routed_scaling;
20286        assert!((glm_wsum - expected).abs() < 1e-15);
20287        assert!(generic_wsum > glm_wsum * 100.0);
20288    }
20289
20290    #[test]
20291    fn resonance_scores_match_formula_and_stable_tie() {
20292        let r = Resonance {
20293            // Three descriptors, hidden=2, one projection row each.
20294            mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0],
20295            u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
20296            k: 1,
20297            bias: vec![1.5, 0.5, 0.0],
20298            shell: Vec::new(),
20299        };
20300        let x = [1.0f32, 1.0];
20301        let mut got = vec![0.0; 3];
20302        r.scores(&x, &mut got);
20303        // Expert 0 and 1 are an exact score tie; the CPU top-1 contract uses
20304        // the lower index.  The values also check d² - (U·d)², not just tie
20305        // ordering.
20306        assert!((got[0] - 0.5).abs() < 1e-6);
20307        assert!((got[1] - 0.5).abs() < 1e-6);
20308        assert!(got[2].abs() < 1e-6);
20309        let best = got
20310            .iter()
20311            .enumerate()
20312            .max_by(|(ia, a), (ib, b)| a.partial_cmp(b).unwrap().then(ib.cmp(ia)))
20313            .map(|(i, _)| i);
20314        assert_eq!(best, Some(0));
20315        assert!(got.iter().all(|v| v.is_finite()));
20316    }
20317
20318    /// The growth shell (spec §2): a grown expert whose reconstruction
20319    /// error lies outside its shell scores −∞, one inside keeps the exact
20320    /// resonance score, trunk rows (`+inf` shell) are bit-identical to the
20321    /// shell-less computation; the process-wide switch disables it.
20322    #[test]
20323    fn resonance_shell_masks_outside_keeps_inside_and_trunk_bits() {
20324        // hidden = 2, rank 1. Experts 0/1 = trunk (shell +inf); 2 and 3 =
20325        // grown, the same descriptor (μ = (0, 1), u = (1, 1)) with shells
20326        // 6.0 and 0.25. At x' = (3, 0): d = (3, −1), d² = 10, proj =
20327        // (3 − 1)² = 4, err = 6 exactly — on the boundary of expert 2's
20328        // shell (kept: the rule is strict `>`), outside expert 3's.
20329        let plain = Resonance {
20330            mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 1.0],
20331            u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0, 1.0],
20332            k: 1,
20333            bias: vec![1.5, 0.5, 0.0, 0.0],
20334            shell: Vec::new(),
20335        };
20336        let shelled = Resonance {
20337            mu: plain.mu.clone(),
20338            u: plain.u.clone(),
20339            k: 1,
20340            bias: plain.bias.clone(),
20341            shell: vec![f32::INFINITY, f32::INFINITY, 6.0, 0.25],
20342        };
20343        assert!(!plain.has_shell());
20344        assert!(shelled.has_shell());
20345        let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
20346        let (mut a, mut b) = (vec![0.0; 4], vec![0.0; 4]);
20347        set_growth_shell(Some(true));
20348        assert!(growth_shell_enabled());
20349        // x = (1, 1): the grown experts reconstruct it exactly (err 0):
20350        // inside both shells, every row the shell-less bits.
20351        let x = [1.0f32, 1.0];
20352        plain.scores(&x, &mut a);
20353        shelled.scores(&x, &mut b);
20354        assert_eq!(bits(&a), bits(&b), "inside every shell: unchanged");
20355        assert!(a[2] == 0.0 && a[3] == 0.0);
20356        // x' = (3, 0): expert 3 → −∞, expert 2 (err == shell) and the
20357        // trunk rows keep their exact bits.
20358        let xo = [3.0f32, 0.0];
20359        plain.scores(&xo, &mut a);
20360        shelled.scores(&xo, &mut b);
20361        assert_eq!(a[2], -6.0);
20362        assert_eq!(a[3], -6.0);
20363        assert_eq!(bits(&a[..3]), bits(&b[..3]), "trunk rows + the boundary row");
20364        assert_eq!(b[3], f32::NEG_INFINITY, "outside the shell: −∞");
20365        assert_eq!(shelled.effective_shell(4), shelled.shell);
20366        // The switch (`CMF_GROWTH_SHELL=off` / `growth-eval --shell off`):
20367        // all +inf, the shell-less bits everywhere.
20368        set_growth_shell(Some(false));
20369        assert!(!growth_shell_enabled());
20370        assert_eq!(shelled.effective_shell(4), vec![f32::INFINITY; 4]);
20371        shelled.scores(&xo, &mut b);
20372        assert_eq!(bits(&a), bits(&b));
20373        set_growth_shell(None);
20374        // A shell vector shorter than the expert count masks nothing
20375        // beyond it (a legacy layer whose tail has no shell).
20376        let short = Resonance {
20377            shell: vec![f32::INFINITY, f32::INFINITY],
20378            ..shelled
20379        };
20380        set_growth_shell(Some(true));
20381        short.scores(&xo, &mut b);
20382        assert_eq!(bits(&a), bits(&b));
20383        set_growth_shell(None);
20384    }
20385
20386    /// `moe_route` with −∞ logits (a grown expert outside its shell):
20387    /// top-1 is the best finite expert with weight exactly 1.0 on both
20388    /// the softmax and the sigmoid path; all −∞ degrades to uniform.
20389    #[test]
20390    fn moe_route_neg_inf_logits_pick_the_best_finite_expert_with_weight_one() {
20391        let zero = || DenseFfn {
20392            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20393            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20394            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20395            act: Act::Silu,
20396            down_t: None,
20397            segs: Vec::new(),
20398        };
20399        let moe = |sigmoid: bool| MoeFfn {
20400            router: QTensor::from_f32(vec![0.0; 8], 4, 2),
20401            experts: vec![zero(), zero(), zero(), zero()],
20402            top_k: 1,
20403            norm_topk_prob: true,
20404            router_sigmoid: sigmoid,
20405            expert_bias: None,
20406            routed_scaling: 1.0,
20407            route_tau: None,
20408            shared: None,
20409            stats: std::cell::RefCell::new(Vec::new()),
20410            act_sq: std::cell::RefCell::new(Vec::new()),
20411            act_rows: std::cell::RefCell::new(Vec::new()),
20412            mask: None,
20413            per_expert_scale: None,
20414            router_input_norm: false,
20415            resonance: None,
20416            grown: Vec::new(),
20417        };
20418        let logits = [-1.0f32, f32::NEG_INFINITY, -0.5, f32::NEG_INFINITY];
20419        for sigmoid in [false, true] {
20420            let m = moe(sigmoid);
20421            let (idx, p, wsum) = moe_route(&logits, &m, None);
20422            assert_eq!(idx, vec![2], "sigmoid {sigmoid}");
20423            assert_eq!(p[1], 0.0);
20424            assert_eq!(p[3], 0.0);
20425            assert!(p[2] > p[0] && p[0] > 0.0);
20426            assert!(p.iter().all(|v| v.is_finite()));
20427            let w = p[2] / wsum;
20428            if sigmoid {
20429                // The sigmoid renorm keeps its reference floor (+1e-6).
20430                assert!((w - 1.0).abs() < 1e-5, "sigmoid: weight {w}");
20431            } else {
20432                assert_eq!(w, 1.0, "softmax: the renormalized top-1 weight is exactly 1");
20433            }
20434            // Masked experts stay masked even when they are the only ones
20435            // "admitted" by an allow-list that covers everything.
20436            let (idx, _, _) = moe_route(&logits, &m, Some(&[true, true, true, true]));
20437            assert_eq!(idx, vec![2]);
20438        }
20439        // A finite expert always beats −∞ whatever the bias / order.
20440        let m = moe(false);
20441        let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, f32::NEG_INFINITY, f32::NEG_INFINITY, -9.0], &m, None);
20442        assert_eq!(idx, vec![3]);
20443        // Every expert at −∞ (cannot happen on a grown file — trunk rows
20444        // have no shell): uniform, finite, lowest index.
20445        let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 4], &m, None);
20446        assert_eq!(idx, vec![0]);
20447        assert!(p.iter().all(|&v| v == 0.25));
20448        assert!(wsum.is_finite() && wsum > 0.0);
20449    }
20450
20451    /// The resonance router (top-1) selects by the raw score as the
20452    /// trainer and the graph do — not by softmax probabilities, where two
20453    /// scores closer than 2^-25 collapse to the same `exp(l − max) = 1.0`
20454    /// and the LOWER index wins a token whose score is strictly smaller.
20455    #[test]
20456    fn resonance_route_picks_the_raw_argmax_not_the_softmax_tie() {
20457        let zero = || DenseFfn {
20458            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20459            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20460            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20461            act: Act::Silu,
20462            down_t: None,
20463            segs: Vec::new(),
20464        };
20465        let moe = |resonant: bool, norm_topk: bool| MoeFfn {
20466            router: QTensor::from_f32(vec![0.0; 6], 3, 2),
20467            experts: vec![zero(), zero(), zero()],
20468            top_k: 1,
20469            norm_topk_prob: norm_topk,
20470            router_sigmoid: false,
20471            expert_bias: None,
20472            routed_scaling: 1.0,
20473            route_tau: None,
20474            shared: None,
20475            stats: std::cell::RefCell::new(Vec::new()),
20476            act_sq: std::cell::RefCell::new(Vec::new()),
20477            act_rows: std::cell::RefCell::new(Vec::new()),
20478            mask: None,
20479            per_expert_scale: None,
20480            router_input_norm: false,
20481            resonance: resonant.then(|| Resonance {
20482                mu: vec![0.0; 6],
20483                u: Vec::new(),
20484                k: 0,
20485                bias: vec![0.0; 3],
20486                shell: Vec::new(),
20487            }),
20488            grown: Vec::new(),
20489        };
20490        // lo = −0.1, hi = the next f32 towards zero: hi − lo = 2^-27 <
20491        // 2^-25, so exp(lo − hi) rounds to exactly 1.0 — a softmax tie.
20492        let lo = -0.1f32;
20493        let hi = f32::from_bits(lo.to_bits() - 1);
20494        assert!(hi > lo && hi - lo < 2f32.powi(-25));
20495        assert_eq!((lo - hi).exp(), 1.0, "the tie this test is about");
20496        // The gated MoE (softmax) path: the tie hands the token to index 0.
20497        let (idx, _, _) = moe_route(&[lo, hi], &moe(false, true), None);
20498        assert_eq!(idx, vec![0], "softmax tie → lower index (the gated contract)");
20499        // The resonance path: the strictly larger raw score wins, weight
20500        // exactly 1.0 with and without norm_topk.
20501        for norm in [true, false] {
20502            let m = moe(true, norm);
20503            let (idx, p, wsum) = moe_route(&[lo, hi, f32::NEG_INFINITY], &m, None);
20504            assert_eq!(idx, vec![1], "norm_topk {norm}");
20505            assert_eq!(p, vec![0.0, 1.0, 0.0]);
20506            assert_eq!(p[1] / wsum, 1.0);
20507            // An exact tie: the first maximum (as `resonance_winner` and
20508            // `embryo_core_route_pick`).
20509            let (idx, _, _) = moe_route(&[hi, hi, lo], &m, None);
20510            assert_eq!(idx, vec![0]);
20511            // `−∞` never wins; the admitted set is honoured.
20512            let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, lo, hi], &m, None);
20513            assert_eq!(idx, vec![2]);
20514            let (idx, p, wsum) = moe_route(&[lo, hi, hi], &m, Some(&[true, false, false]));
20515            assert_eq!(idx, vec![0]);
20516            assert_eq!(p[0] / wsum, 1.0);
20517            // Every admitted expert at −∞: the generic path's uniform
20518            // fallback (lowest index, finite weights).
20519            let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 3], &m, None);
20520            assert_eq!(idx, vec![0]);
20521            assert!(p.iter().all(|v| v.is_finite()) && wsum.is_finite() && wsum > 0.0);
20522        }
20523    }
20524}