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    /// Final-logit soft-capping C: logits = C·tanh(logits/C) (Gemma-4).
388    pub final_softcap: Option<f32>,
389    /// Cortiq Embryo hierarchical head: cluster matrix [C, hidden]. The
390    /// flat logits h·Eᵀ are turned into the two-level log-probabilities
391    /// log softmax_c(h·Cᵀ)[c(v)] + log softmax_{s∈c(v)}(h·E_c(v)ᵀ)[v].
392    pub head_clusters: Option<std::sync::Arc<Vec<f32>>>,
393    /// Gemma-2 attention-logit soft-capping (0.0 = off).
394    pub attn_softcap: f32,
395    /// Compute per-token confidence (a full-vocab softmax each
396    /// token). On by default; `bench --core` turns it off to match
397    /// llama-bench's core timing.
398    confidence_on: bool,
399    /// Test-only one-shot forward failure, scoped to this pipeline so
400    /// parallel scoring tests cannot consume one another's injection.
401    #[cfg(test)]
402    nll_test_fail_at: Option<usize>,
403    /// Test-only route override; avoids mutating the process-wide
404    /// `CMF_PREFILL` environment variable while forcing the serial path.
405    #[cfg(test)]
406    nll_test_force_serial: bool,
407}
408
409#[cfg(target_os = "macos")]
410impl Drop for Pipeline {
411    fn drop(&mut self) {
412        // the async replay writes into `kv_cache` Vecs about to be freed
413        let _ = crate::gpu_metal::wait_replay();
414        crate::gpu::kv_mirror_drop(self.graph_kv_id);
415    }
416}
417
418#[cfg(not(target_os = "macos"))]
419impl Drop for Pipeline {
420    fn drop(&mut self) {
421        // The wgpu resident Embryo graph owns recurrent/KV buffers keyed by
422        // this pipeline's sequence id.  Release that sequence image when a
423        // pooled pipeline is dropped; model weights stay cached for reuse.
424        crate::gpu::graph_kv_reset(self.graph_kv_id);
425    }
426}
427
428/// Model weights. Matrices are `QTensor` (owned f32 for small models
429/// and tests — bit-identical to the historical paths — or quantized
430/// bytes zero-copy from the CMF mmap for big models). 1-D norms are
431/// always small and stay f32.
432pub struct PipelineWeights {
433    /// Embedding table: [vocab_size, hidden_size]
434    pub embed_tokens: QTensor,
435    /// Per-layer weights
436    pub layers: Vec<LayerWeights>,
437    /// LM head: [vocab_size, hidden_size]
438    pub lm_head: QTensor,
439    /// Final norm: [hidden_size]
440    pub final_norm: Vec<f32>,
441}
442
443/// One transformer layer: shared norms + MLP, attention by kind.
444pub struct LayerWeights {
445    pub input_norm: Vec<f32>,
446    /// The pre-FFN norm (`post_attention_layernorm` classically;
447    /// `pre_feedforward_layernorm` on Gemma-2/3 sandwich layers).
448    pub post_norm: Vec<f32>,
449    /// Gemma-2/3 sandwich: norm applied to the ATTENTION OUTPUT before
450    /// its residual add (`post_attention_layernorm` there).
451    pub attn_out_norm: Option<Vec<f32>>,
452    /// Gemma-4: the whole layer output is multiplied by this scalar.
453    pub layer_scale: Option<f32>,
454    /// Gemma-2/3 sandwich: norm applied to the FFN OUTPUT before its
455    /// residual add (`post_feedforward_layernorm`).
456    pub ffn_out_norm: Option<Vec<f32>>,
457    pub ffn: FfnKind,
458    pub attn: AttnKind,
459}
460
461/// FFN gate activation: SiLU (SwiGLU family) or tanh-GELU (Gemma's
462/// GeGLU). A property of the model, carried on every FFN triple.
463#[derive(Clone, Copy, PartialEq, Debug, Default)]
464pub enum Act {
465    #[default]
466    Silu,
467    GeluTanh,
468    /// Kimi-K3 SituAndMul: BOTH halves transform —
469    /// a = β·tanh(g/β)·σ(g), up' = linβ·tanh(u/linβ) (linβ>0), out = a·up'.
470    Situ {
471        beta: f32,
472        linear_beta: f32,
473    },
474}
475
476impl Act {
477    pub fn from_arch(name: &str) -> Self {
478        if name == "gelu_tanh" {
479            Self::GeluTanh
480        } else {
481            Self::Silu
482        }
483    }
484
485    /// Arch-driven constructor (activation name + situ betas).
486    pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
487        match arch.hidden_act.as_str() {
488            "situ" => Self::Situ {
489                beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
490                linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
491            },
492            other => Self::from_arch(other),
493        }
494    }
495
496    #[inline]
497    pub fn apply(self, x: f32) -> f32 {
498        match self {
499            Self::Silu => inference::silu(x),
500            Self::GeluTanh => inference::gelu_tanh(x),
501            Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
502        }
503    }
504
505    /// Gated combine — the FFN contract. Situ transforms the UP half
506    /// too, so callers must use this instead of apply(g)·u.
507    #[inline]
508    pub fn combine(self, g: f32, u: f32) -> f32 {
509        match self {
510            Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
511                self.apply(g) * (linear_beta * (u / linear_beta).tanh())
512            }
513            _ => self.apply(g) * u,
514        }
515    }
516}
517
518/// Dense gated triple — the FFN of a dense layer or of one expert.
519pub struct DenseFfn {
520    pub gate_proj: QTensor,
521    pub up_proj: QTensor,
522    pub down_proj: QTensor,
523    /// Gate activation (SiLU default; Gemma: tanh-GELU).
524    pub act: Act,
525    /// `down_proj` stored transposed (`[inter, hidden]`), when the file
526    /// carries it. Only the per-token sparse path reads it: a neuron's
527    /// down weights are a contiguous ROW there, so the token's chosen
528    /// neurons are the only bytes touched. `None` = the ordinary layout,
529    /// and the sparse path stays off.
530    pub down_t: Option<QTensor>,
531    /// Task tubes (spec: defragged task-conditional width). The three
532    /// matrices above are the CORE — the neurons every task computes;
533    /// each tube is an independently quantized slice of the SAME layer
534    /// holding the neurons only some tasks need. A tube is a normal
535    /// tensor triple, so every kernel runs it unchanged, and the bytes
536    /// of an inactive tube are never read. Empty = ordinary dense FFN.
537    pub segs: Vec<FfnSeg>,
538}
539
540/// One task tube: a contiguous slice of a layer's FFN neurons, stored
541/// as its own `[w, hidden]` / `[hidden, w]` triple. `start` is the
542/// neuron's index in the layer's FULL space (core first, then tubes in
543/// order) — the bit a task mask sets to switch this tube on.
544pub struct FfnSeg {
545    pub gate: QTensor,
546    pub up: QTensor,
547    pub down: QTensor,
548    pub start: usize,
549    pub width: usize,
550}
551
552/// FFN operator of a layer, decided by tensor presence at load time
553/// (router `mlp.gate.weight` in the directory = MoE layer).
554pub enum FfnKind {
555    Dense(DenseFfn),
556    /// Mixture-of-Experts (Qwen2-MoE / Qwen3-MoE): softmax over ALL
557    /// expert logits → top-k, optional renorm; experts stay quantized
558    /// in mmap — only the selected ones are touched per token.
559    Moe(MoeFfn),
560    /// Gemma-4 MoE: a dense MLP branch AND a routed-expert branch in
561    /// the SAME layer, each with its own norm sandwich. The dense
562    /// branch reads the pre-FFN-normed input; the expert branch (and
563    /// the router) read the RAW residual through `pre_norm_2`:
564    ///   d = post_norm_1(dense(x̂));  m = post_norm_2(Σwₑ·FFNₑ(pre_norm_2(h)))
565    ///   ffn_out = d + m   (the caller's ffn_out_norm + residual follow)
566    DenseMoe(Box<DenseMoeFfn>),
567}
568
569/// Gemma-4 dual-branch FFN (see `FfnKind::DenseMoe`).
570pub struct DenseMoeFfn {
571    pub dense: DenseFfn,
572    pub moe: MoeFfn,
573    /// post_feedforward_layernorm_1 — dense-branch output norm.
574    pub post_norm_1: Vec<f32>,
575    /// pre_feedforward_layernorm_2 — expert-branch input norm (applied
576    /// to the RAW residual, not the pre-FFN-normed activation).
577    pub pre_norm_2: Vec<f32>,
578    /// post_feedforward_layernorm_2 — expert-branch output norm.
579    pub post_norm_2: Vec<f32>,
580}
581
582pub struct MoeFfn {
583    /// Router `mlp.gate.weight` [num_experts, hidden].
584    pub router: QTensor,
585    pub experts: Vec<DenseFfn>,
586    pub top_k: usize,
587    pub norm_topk_prob: bool,
588    /// Router scores per-expert with a sigmoid (LFM2-MoE / DeepSeek-V3
589    /// `noaux_tc`) instead of a softmax over all experts (Qwen).
590    pub router_sigmoid: bool,
591    /// Per-expert selection bias `mlp.expert_bias` [num_experts]
592    /// (LFM2-MoE): added to the sigmoid scores for the top-k CHOICE only;
593    /// the gathered weights use the unbiased scores. None = no bias.
594    pub expert_bias: Option<Vec<f32>>,
595    /// Top-k weights are multiplied by this after the optional renorm
596    /// (LFM2-MoE `routed_scaling_factor`; 1.0 = off).
597    pub routed_scaling: f32,
598    /// Adaptive routing (CMF_MOE_TAU, opt-in): keep the smallest
599    /// prefix of the top-k whose renormalized mass reaches τ —
600    /// confident tokens touch 1–2 experts, flat ones keep all k.
601    /// MoE decode is memory-bound, so skipped experts are skipped
602    /// weight traffic. None = classic fixed top-k (bit-identical).
603    pub route_tau: Option<f32>,
604    /// Always-on shared expert. Qwen2-MoE carries an additional sigmoid
605    /// gate; Laguna adds the shared expert unconditionally (`None`).
606    pub shared: Option<(DenseFfn, Option<QTensor>)>,
607    /// Expert-selection counters (truncated Fisher B-field of claim 12:
608    /// routing frequency during calibration). Filled by every forward,
609    /// read by the CLI via CMF_MOE_STATS. RefCell: decode is single-threaded.
610    pub stats: std::cell::RefCell<Vec<u64>>,
611    /// Per-CHANNEL sum of squares of this FFN's input, accumulated over a
612    /// calibration run (`CMF_RMS_TRACE`). These are the RMS activation
613    /// traces AWNP needs: raw weight magnitude says every channel matters
614    /// equally, and the question AWNP asks is whether the ACTIVATIONS
615    /// disagree. Off unless the env var is set — an f64 add per channel
616    /// per token is cheap, but not free.
617    pub act_sq: std::cell::RefCell<Vec<f64>>,
618    /// Raw FFN-input rows captured for the layers named by `CMF_ACT_DUMP`
619    /// (`"9,19"`). AWNP is nullspace PROJECTION: after dropping channels the
620    /// survivors are refitted to absorb what was removed, and how much they
621    /// can absorb depends on the activation COVARIANCE, not on per-channel
622    /// RMS. Per-channel numbers can only bound the cost from above.
623    pub act_rows: std::cell::RefCell<Vec<f32>>,
624    /// Task mask over routed experts (DTG-MA over MoE, claim-12 B-field
625    /// applied): `false` experts are excluded from selection, the
626    /// softmax renormalizes over the allowed set. Built by the loader
627    /// from CMF_MOE_MASK=<stats.json> + CMF_MOE_MASK_COVER. None = all.
628    pub mask: Option<Vec<bool>>,
629    /// Gemma-4: per-expert weight scale applied AFTER the top-k renorm
630    /// (`router.per_expert_scale`). None = 1.0 everywhere.
631    pub per_expert_scale: Option<Vec<f32>>,
632    /// Gemma-4: the router reads a SCALE-LESS rms-norm of its input
633    /// (the constant gain router.scale·√hidden is folded into the
634    /// router weights at convert time).
635    pub router_input_norm: bool,
636    /// Cortiq Embryo: resonance routing (P1) — the "logits" are
637    /// bias_e − ‖(x−μ_e) − U_eᵀU_e(x−μ_e)‖², argmax = the expert whose
638    /// descriptor reconstructs the input best. `router` is a placeholder.
639    pub resonance: Option<Resonance>,
640    /// Growth records (`kind = "expert_append"`, spec §9.5.1) mounted
641    /// behind the trunk experts, one entry per grown expert IN THE ORDER
642    /// they sit in `experts` (the tail `experts[experts.len() - grown.len()..]`).
643    /// The trunk keeps `experts.len() - grown.len()` experts. Empty on a
644    /// gated MoE, on a file without records and under `CMF_GROWTH=off`.
645    pub grown: Vec<GrownExpert>,
646}
647
648/// One grown expert of an `expert_append` record as the loader mounted it.
649#[derive(Debug, Clone, PartialEq, Eq)]
650pub struct GrownExpert {
651    /// The record's skill id.
652    pub record: String,
653    /// Position of the record in `header.skills`.
654    pub record_index: usize,
655    pub layer: usize,
656    /// The index the tensor name declares (`experts.{e}` — the chain rule
657    /// of the format); the executed position may be smaller when an
658    /// earlier record of the layer is not mounted.
659    pub expert: usize,
660}
661
662/// Per-expert resonance descriptors of one MoE layer (`mlp.desc.*`).
663pub struct Resonance {
664    /// [E, hidden]
665    pub mu: Vec<f32>,
666    /// [E, k, hidden] orthonormal directions (k may be 0)
667    pub u: Vec<f32>,
668    pub k: usize,
669    /// [E] selection bias (loss-free balancing, trained online)
670    pub bias: Vec<f32>,
671    /// [E] reconstruction-error shell: an expert whose error
672    /// `‖(x−μ)⊥U‖² = d² − proj` exceeds its shell scores `−∞` (it never
673    /// wins). `+inf` = no shell — every trunk expert; a grown expert
674    /// (`expert_append`) carries the finite `desc.shell` its record stores.
675    /// Empty = no shell anywhere (legacy constructors).
676    pub shell: Vec<f32>,
677}
678
679/// `CMF_GROWTH_SHELL` state: 0 = not yet read from the environment, 1 = on,
680/// 2 = off. Process-wide, like the environment it mirrors.
681static GROWTH_SHELL: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
682
683/// Is the growth shell applied (`Resonance::scores` −∞ rule, the resident
684/// graph's packed shell)? `CMF_GROWTH_SHELL=off` disables it for
685/// measurement; [`set_growth_shell`] overrides the environment in-process
686/// (`growth-eval --shell`). Default: on.
687pub fn growth_shell_enabled() -> bool {
688    use std::sync::atomic::Ordering;
689    match GROWTH_SHELL.load(Ordering::Relaxed) {
690        1 => true,
691        2 => false,
692        _ => {
693            let off = std::env::var("CMF_GROWTH_SHELL")
694                .map(|v| v.eq_ignore_ascii_case("off") || v == "0")
695                .unwrap_or(false);
696            GROWTH_SHELL.store(if off { 2 } else { 1 }, Ordering::Relaxed);
697            !off
698        }
699    }
700}
701
702/// Switch the growth shell on/off for this process (`None` = re-read
703/// `CMF_GROWTH_SHELL` on the next query). A pipeline packed into the
704/// resident graph BEFORE the switch keeps the shell it was packed with —
705/// build a new pipeline after switching.
706pub fn set_growth_shell(on: Option<bool>) {
707    GROWTH_SHELL.store(
708        match on {
709            Some(true) => 1,
710            Some(false) => 2,
711            None => 0,
712        },
713        std::sync::atomic::Ordering::Relaxed,
714    );
715}
716
717impl Resonance {
718    /// Does any expert carry a finite shell (a mounted growth record)?
719    pub fn has_shell(&self) -> bool {
720        self.shell.iter().any(|s| s.is_finite())
721    }
722
723    /// The shell the runtime applies right now: the stored one, or all
724    /// `+inf` when the shell is switched off (`CMF_GROWTH_SHELL=off`).
725    pub fn effective_shell(&self, ne: usize) -> Vec<f32> {
726        let mut out = vec![f32::INFINITY; ne];
727        if growth_shell_enabled() {
728            for (o, s) in out.iter_mut().zip(&self.shell) {
729                *o = *s;
730            }
731        }
732        out
733    }
734
735    /// Routing scores for one input row (higher = better). A grown
736    /// expert whose reconstruction error lies outside its shell gets
737    /// `−∞` (unless the shell is switched off); trunk rows are the exact
738    /// bit pattern they were before growth.
739    pub fn scores(&self, x: &[f32], out: &mut [f32]) {
740        let h = x.len();
741        let ne = out.len();
742        let shell_on = growth_shell_enabled() && !self.shell.is_empty();
743        for e in 0..ne {
744            let mu = &self.mu[e * h..(e + 1) * h];
745            let mut d2 = 0.0f32;
746            for j in 0..h {
747                let d = x[j] - mu[j];
748                d2 += d * d;
749            }
750            let mut proj = 0.0f32;
751            for i in 0..self.k {
752                let u = &self.u[(e * self.k + i) * h..(e * self.k + i + 1) * h];
753                let mut p = 0.0f32;
754                for j in 0..h {
755                    p += (x[j] - mu[j]) * u[j];
756                }
757                proj += p * p;
758            }
759            let err = d2 - proj;
760            out[e] = self.bias.get(e).copied().unwrap_or(0.0) - err;
761            if shell_on && err > self.shell.get(e).copied().unwrap_or(f32::INFINITY) {
762                out[e] = f32::NEG_INFINITY;
763            }
764        }
765    }
766}
767
768/// Attention operator of a layer. Extension point: new operators are
769/// new variants here + a forward in their own module.
770pub enum AttnKind {
771    /// GQA softmax attention (+ optional Qwen3.5 qk-norm / output gate).
772    Full {
773        wq: QTensor,
774        wk: QTensor,
775        wv: QTensor,
776        wo: QTensor,
777        q_norm: Option<Vec<f32>>,
778        k_norm: Option<Vec<f32>>,
779        output_gate: bool,
780        /// Laguna: a separate softplus projection applied to the attention
781        /// output before O. The bool means one scalar per head (broadcast
782        /// across head_dim); false means one scalar per element.
783        softplus_gate: Option<(QTensor, bool)>,
784        /// Qwen2-family projection biases (q, k, v).
785        bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
786    },
787    /// Canonical linear core (VMF phase attention).
788    Linear(VmfPhaseWeights),
789    /// Faithful vendor linear operator (Qwen3.5 GatedDeltaNet).
790    LinearGdn(GdnWeights),
791    /// LFM2 gated short-convolution mixer (no KV cache; conv ring state
792    /// lives in the layer's `linear_state`).
793    ShortConv(ShortConvWeights),
794    /// DeepSeek-V2 Multi-head Latent Attention. v1 executes it as
795    /// expand-to-MHA: the latent is projected per token, K/V expand to
796    /// every head and live in the ordinary cache (K head layout
797    /// [rope | nope] so the standard partial rotary covers the shared
798    /// rope key; V rows are zero-padded to the K head_dim and the pad
799    /// is sliced off before O). Latent-resident cache is a later
800    /// optimization, not a semantic change.
801    Mla(Box<MlaWeights>),
802    /// Kimi Delta Attention (Kimi Linear / Kimi-K3): per-channel decayed
803    /// delta rule, separate q/k/v short convs, sigmoid-gated output norm.
804    /// State lives in the layer's `linear_state` (no KV cache).
805    Kda(Box<crate::linear_core::KdaWeights>),
806    /// Natively bounded softmax anchor `swa_sink_v1` (Embryo-O1): ring of
807    /// the last W raw keys with relative RoPE + trained NoPE sinks, one
808    /// softmax. State is the fixed-size ring in `LayerKvCache::bounded`;
809    /// nothing is stored per position (see `crate::bounded`).
810    Bounded(Box<crate::bounded::BoundedWeights>),
811}
812
813/// DeepSeek-V2 MLA projections (see `AttnKind::Mla`).
814pub struct MlaWeights {
815    /// `[nh·(rope+nope), hidden]` (or `[…, q_lora]` when compressed) —
816    /// the converter permutes each head rope-first so rotary_dim =
817    /// qk_rope works unchanged.
818    pub q_proj: QTensor,
819    /// Compressed q (K3/V3 class): x → q_a `[q_lora, hidden]` →
820    /// rms(q_a_norm) → q_proj (= q_b). None = direct q (V2-Lite).
821    pub q_a: Option<QTensor>,
822    pub q_a_norm: Option<Vec<f32>>,
823    /// `kv_a_proj_with_mqa` `[lora + rope, hidden]` (latent first).
824    pub kv_a: QTensor,
825    /// RMS-norm weights over the latent (`kv_a_layernorm`, [lora]).
826    pub kv_a_norm: Vec<f32>,
827    /// `[nh·(nope+v), lora]` — per head [k_nope | v].
828    pub kv_b: QTensor,
829    /// `[hidden, nh·v]`.
830    pub o_proj: QTensor,
831    pub nh: usize,
832    pub qk_rope: usize,
833    pub qk_nope: usize,
834    pub v_dim: usize,
835    pub lora: usize,
836    /// Softmax scale (1/√(rope+nope), YaRN-mscale-corrected at load).
837    pub scale: f32,
838    /// Kimi Linear NoPE: skip the rotary entirely (layout unchanged).
839    pub nope: bool,
840}
841
842/// Multi-token-prediction head (DeepSeek/Qwen style, spec §2.1):
843/// `x = eh_proj·[enorm(embed(next)); hnorm(hidden)]` → one transformer
844/// block over its own KV → shared lm_head. Drafts the token after next;
845/// the main model verifies, so output is exact — MTP only buys speed.
846pub struct MtpModule {
847    pub enorm: Vec<f32>,
848    pub hnorm: Vec<f32>,
849    /// [hidden, 2·hidden]
850    pub eh_proj: QTensor,
851    pub layer: LayerWeights,
852    pub final_norm: Vec<f32>,
853    pub kv: crate::kv_cache::LayerKvCache,
854}
855
856/// A Metal verify graph after its sync: what the commit needs — the
857/// graph (per-layer replay scratch), the GDN layers in encode order (their
858/// CPU states receive the replay), and the attention layers with the CPU
859/// row count they were encoded against (the accepted rows are pulled from
860/// the mirror from there).
861/// One item of the Metal rows-graph plan.
862#[cfg(target_os = "macos")]
863enum MetalRowsItem<'a> {
864    Gdn {
865        run: Vec<crate::gpu_metal::GdnGpuLayer<'a>>,
866        first: usize,
867    },
868    Attn {
869        l: crate::gpu_metal::AttnGpuLayer<'a>,
870        li: usize,
871        q_norm: Option<&'a [f32]>,
872        k_norm: Option<&'a [f32]>,
873        output_gate: bool,
874    },
875}
876
877#[cfg(target_os = "macos")]
878struct MetalVerifyPending {
879    graph: crate::gpu_metal::VerifyGraph,
880    gdn_layers: Vec<usize>,
881    attn_layers: Vec<(usize, usize)>,
882}
883
884/// A round's batched MTP warm-up, submitted but not yet waited
885/// (`mtp_warm_batch_submit` → `mtp_warm_batch_finish`): the trunk commit's
886/// GDN replay is queued between the two.
887#[cfg(target_os = "macos")]
888struct MetalWarmPending {
889    graph: crate::gpu_metal::VerifyGraph,
890    cpu_stored: usize,
891    b: usize,
892}
893
894#[cfg(target_os = "macos")]
895enum MetalRowsRun {
896    /// Capability/preflight refusal before a command buffer was committed.
897    Declined,
898    /// A graph was admitted and then failed; callers must clear the sequence
899    /// rather than replaying it through CPU/serial state.
900    Failed,
901    Completed(MetalVerifyPending),
902}
903
904#[cfg(target_os = "macos")]
905enum MetalPrefillOutcome {
906    Declined,
907    Failed,
908    Completed(Vec<f32>),
909}
910
911#[cfg(target_os = "macos")]
912enum MetalBatchNllOutcome {
913    Declined,
914    Failed(String),
915    Completed(f64, usize),
916}
917
918/// The speculation trial's phases (see the decode loop): four timed
919/// speculative rounds, eight timed plain tokens, then the faster arm
920/// until a re-check.
921#[derive(Clone, Copy)]
922enum SpecTrial {
923    Spec {
924        t0: std::time::Instant,
925        gen0: usize,
926        rounds: usize,
927    },
928    Plain {
929        t0: std::time::Instant,
930        gen0: usize,
931    },
932    Decided {
933        spec: bool,
934        recheck_at: usize,
935    },
936}
937
938/// `CMF_GRAPH_SPEC_TIME`: 0 = off, 1 = one line per speculative round
939/// plus the host stamps of any OUTLIER round (wall > 1.4× the running
940/// median), 2 = the host stamps of every round.
941pub(crate) fn spec_time_level() -> u8 {
942    static L: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
943    *L.get_or_init(|| match std::env::var("CMF_GRAPH_SPEC_TIME") {
944        Ok(v) => v.trim().parse::<u8>().map(|n| n.max(1)).unwrap_or(1),
945        Err(_) => 0,
946    })
947}
948
949/// The round's host stamps: `spec_stamp(name)` records the time since
950/// the previous stamp (the section that just ended) — from anywhere on
951/// the round's call chain (the Metal verify, the draft step, the commit),
952/// no plumbing. Off (a single atomic load) unless `CMF_GRAPH_SPEC_TIME`
953/// is set; one decode thread at a time is assumed (diagnostics).
954struct SpecStampLog {
955    t_last: std::time::Instant,
956    items: Vec<(&'static str, f32)>,
957}
958
959static SPEC_STAMPS: std::sync::Mutex<Option<SpecStampLog>> = std::sync::Mutex::new(None);
960
961pub(crate) fn spec_stamp(name: &'static str) {
962    if spec_time_level() == 0 {
963        return;
964    }
965    if let Ok(mut g) = SPEC_STAMPS.lock() {
966        if let Some(log) = g.as_mut() {
967            let now = std::time::Instant::now();
968            log.items
969                .push((name, (now - log.t_last).as_secs_f32() * 1e3));
970            log.t_last = now;
971        }
972    }
973}
974
975fn spec_stamps_begin() {
976    if spec_time_level() == 0 {
977        return;
978    }
979    if let Ok(mut g) = SPEC_STAMPS.lock() {
980        *g = Some(SpecStampLog {
981            t_last: std::time::Instant::now(),
982            items: Vec::with_capacity(64),
983        });
984    }
985}
986
987fn spec_stamps_take() -> Vec<(&'static str, f32)> {
988    SPEC_STAMPS
989        .lock()
990        .ok()
991        .and_then(|mut g| g.take())
992        .map(|l| l.items)
993        .unwrap_or_default()
994}
995
996/// One line: every stamp name in first-seen order with its total over the
997/// round and, when it fired more than once (the draft steps), the count.
998fn spec_stamps_format(items: &[(&'static str, f32)]) -> String {
999    let mut agg: Vec<(&'static str, f32, u32)> = Vec::with_capacity(items.len());
1000    for &(n, ms) in items {
1001        match agg.iter_mut().find(|e| e.0 == n) {
1002            Some(e) => {
1003                e.1 += ms;
1004                e.2 += 1;
1005            }
1006            None => agg.push((n, ms, 1)),
1007        }
1008    }
1009    let mut s = String::with_capacity(agg.len() * 16);
1010    for (n, ms, k) in agg {
1011        if k > 1 {
1012            s.push_str(&format!("{n} {ms:.1}/{k} "));
1013        } else {
1014            s.push_str(&format!("{n} {ms:.1} "));
1015        }
1016    }
1017    s
1018}
1019
1020/// The speculation monitor: exponential averages of a round's wall time
1021/// and of the tokens it produced, and the plain token's wall time — the
1022/// three numbers the keep/stop rule needs. A round pays when
1023/// `tokens_per_round · plain_ms > round_ms · 1.03`. The one-shot trial
1024/// (four rounds against eight tokens) mis-called prose: the first rounds
1025/// after a prompt are formulaic and accept well, the body does not (an
1026/// essay measured 39 against a plain 44.8 with the trial saying
1027/// "speculate"), so the rule now runs on EVERY round and stops after four
1028/// consecutive losing rounds; a stopped speculation is retried 128 tokens
1029/// later.
1030///
1031/// Native Metal (`metal: true`) does not pay the eight plain tokens up
1032/// front: on the 27B a plain token is ~150 ms, so the trial alone cost
1033/// ~1.2 s of every answer. There the plain phase is (a) skipped while the
1034/// rounds land at least `SPEC_PROXY_TOKENS` tokens each — a k=7 round on
1035/// Metal costs ~1.9 plain tokens (286 against 148 ms measured on the M4),
1036/// so 3.5 tokens/round cannot lose on any Metal round/plain ratio seen —
1037/// and (b) otherwise bounded to the fewest tokens that time it: two, or
1038/// as many as fit in `SPEC_PLAIN_MIN_MS` (a 150-ms token measures itself;
1039/// a 10-ms one needs the eight). The keep/stop rule itself is unchanged:
1040/// the moment a plain rate exists, it decides.
1041#[derive(Default, Clone, Copy)]
1042struct SpecMon {
1043    round_ms: f64,
1044    tokens: f64,
1045    plain_ms: f64,
1046    n: u32,
1047    fails: u32,
1048    metal: bool,
1049}
1050
1051/// Tokens per round at or above which a Metal round pays without a plain
1052/// measurement (see `SpecMon`).
1053const SPEC_PROXY_TOKENS: f64 = 3.5;
1054/// The Metal plain phase: at least two tokens, and more until this much
1055/// wall time has been timed (up to the eight the other backends time).
1056const SPEC_PLAIN_MIN_MS: f64 = 200.0;
1057
1058impl SpecMon {
1059    fn round(&mut self, dt_ms: f64, produced: usize) {
1060        self.n += 1;
1061        if self.n == 1 {
1062            return; // round 1 pays the batch scratch and the draft mirror
1063        }
1064        let a = if self.n == 2 { 1.0 } else { 0.3 };
1065        self.round_ms += a * (dt_ms - self.round_ms);
1066        self.tokens += a * (produced as f64 - self.tokens);
1067    }
1068    fn pays(&self) -> bool {
1069        if self.plain_ms > 0.0 {
1070            self.tokens * self.plain_ms > self.round_ms * 1.03
1071        } else {
1072            self.metal && self.tokens >= SPEC_PROXY_TOKENS
1073        }
1074    }
1075    /// Has the plain phase timed enough tokens to decide?
1076    fn plain_done(&self, t0: std::time::Instant, gen0: usize, generated: usize) -> bool {
1077        let n = generated.saturating_sub(gen0);
1078        if n >= 8 {
1079            return true;
1080        }
1081        self.metal && n >= 2 && t0.elapsed().as_secs_f64() * 1e3 >= SPEC_PLAIN_MIN_MS
1082    }
1083}
1084
1085/// Ids of the consumed prefix the bounded reuse key remembers literally
1086/// (the rest is covered by the rolling hash).
1087pub const KV_PREFIX_TAIL: usize = 128;
1088
1089/// Bounded prefix-reuse key: how many ids the cache holds, a rolling
1090/// hash of ALL of them and the last [`KV_PREFIX_TAIL`] ids literally.
1091/// Answers "does this prompt strictly extend what is cached" exactly
1092/// (hash over the whole consumed prefix + literal tail) without keeping
1093/// the dialogue — the record is a fixed size whatever the session length.
1094#[derive(Debug, Clone, Default)]
1095pub struct KvPrefix {
1096    len: usize,
1097    hash: u64,
1098    tail: Vec<u32>,
1099    /// Owner of the state this key describes: the resident device graph
1100    /// (true) or the host. A turn continues the prefix only on its owner
1101    /// (R4: a host continuation of a device sequence reads an empty host
1102    /// state; a device continuation of a host sequence has no image).
1103    device: bool,
1104}
1105
1106impl KvPrefix {
1107    #[inline]
1108    fn fold(mut h: u64, ids: &[u32]) -> u64 {
1109        for &id in ids {
1110            h ^= id as u64;
1111            h = h.wrapping_mul(0x100000001b3);
1112            h ^= h >> 29;
1113        }
1114        h
1115    }
1116
1117    pub fn clear(&mut self) {
1118        self.len = 0;
1119        self.hash = 0xcbf29ce484222325;
1120        self.tail.clear();
1121        self.device = false;
1122    }
1123
1124    /// Was the prefix built on the resident device graph?
1125    pub fn on_device(&self) -> bool {
1126        self.device
1127    }
1128
1129    /// Tag the owner of the recorded prefix.
1130    pub fn set_on_device(&mut self, device: bool) {
1131        self.device = device;
1132    }
1133
1134    /// Ids the cache holds (the forwarded prefix).
1135    pub fn len(&self) -> usize {
1136        self.len
1137    }
1138
1139    pub fn is_empty(&self) -> bool {
1140        self.len == 0
1141    }
1142
1143    /// Literal tail currently kept (≤ `KV_PREFIX_TAIL`).
1144    pub fn tail_len(&self) -> usize {
1145        self.tail.len()
1146    }
1147
1148    /// Replace the key with `ids` (a fresh sequence).
1149    pub fn set(&mut self, ids: &[u32]) {
1150        self.clear();
1151        self.extend(ids);
1152    }
1153
1154    /// Append `more` to the consumed prefix (an extension-only turn).
1155    pub fn extend(&mut self, more: &[u32]) {
1156        if self.len == 0 && self.hash == 0 {
1157            self.hash = 0xcbf29ce484222325;
1158        }
1159        self.hash = Self::fold(self.hash, more);
1160        self.len += more.len();
1161        if more.len() >= KV_PREFIX_TAIL {
1162            self.tail.clear();
1163            self.tail.extend_from_slice(&more[more.len() - KV_PREFIX_TAIL..]);
1164        } else {
1165            let drop = (self.tail.len() + more.len()).saturating_sub(KV_PREFIX_TAIL);
1166            self.tail.drain(..drop);
1167            self.tail.extend_from_slice(more);
1168        }
1169    }
1170
1171    /// Cached positions when `ids` strictly extends the consumed prefix,
1172    /// 0 otherwise. The tail is compared literally first (cheap), then
1173    /// the hash over the whole prefix must agree.
1174    pub fn extension(&self, ids: &[u32]) -> usize {
1175        if self.len == 0 || ids.len() <= self.len {
1176            return 0;
1177        }
1178        let t = self.tail.len();
1179        if ids[self.len - t..self.len] != self.tail[..] {
1180            return 0;
1181        }
1182        if Self::fold(0xcbf29ce484222325, &ids[..self.len]) != self.hash {
1183            return 0;
1184        }
1185        self.len
1186    }
1187}
1188
1189/// Result of a generation call.
1190pub struct GenerateResult {
1191    pub text: String,
1192    pub token_ids: Vec<u32>,
1193    pub prompt_tokens: usize,
1194    pub tokens_generated: usize,
1195    pub finish_reason: String,
1196    /// Speculative-decode stats (0/0 when MTP is absent or inactive).
1197    pub mtp_drafted: usize,
1198    pub mtp_accepted: usize,
1199    /// Per-generated-token confidence = softmax probability of the token
1200    /// that was actually emitted (softmax probability on the chosen state). High =
1201    /// the model was sure; low = it was guessing. Same length as the
1202    /// generated slice of `token_ids`.
1203    pub token_confidence: Vec<f32>,
1204    /// Structured per-token telemetry (B4 channel). Empty unless
1205    /// `set_trace(true)`; otherwise same length as the generated slice.
1206    pub traces: Vec<TokenTrace>,
1207}
1208
1209/// One row of the structured telemetry trace (B4): the model's internal
1210/// routing state at the moment a token was emitted. Every field is a
1211/// quantity the runtime already computes — nothing is inferred or
1212/// estimated (anti-principle: only measured bytes).
1213#[derive(Clone, Debug)]
1214pub struct TokenTrace {
1215    /// 0-based index within the generated slice.
1216    pub t: usize,
1217    /// The emitted token id.
1218    pub token_id: u32,
1219    /// Softmax probability on the emitted token — how sure the model was.
1220    pub confidence: f32,
1221    /// Skill in force while this token was generated (None = backbone).
1222    pub active_skill: Option<String>,
1223    /// Recon error E = ‖r−BBᵀr‖²/‖φ‖² at the last routing eval — coherence
1224    /// with the active skill's subspace (low = coherent). None = no router
1225    /// or not yet evaluated.
1226    pub recon: Option<f32>,
1227    /// The router changed the active skill right after this token (a
1228    /// domain boundary crossed under the hysteresis barrier).
1229    pub switched: bool,
1230}
1231
1232/// Calibrated softmax probability of `id` under `logits` (the confidence on
1233/// the emitted token) — the confidence signal, cheap from logits already
1234/// computed for sampling. `temp` is the calibration temperature (B1):
1235/// softmax(logits / temp); 1.0 = raw.
1236#[cfg_attr(not(test), allow(dead_code))]
1237fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
1238    let t = if temp > 1e-3 { temp } else { 1.0 };
1239    let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1240    let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
1241    if sum > 0.0 {
1242        (((logits[id as usize] - max) / t).exp()) / sum
1243    } else {
1244        0.0
1245    }
1246}
1247
1248/// prefill-GEMM enabled? (CMF_PREFILL=seq — emergency fallback to the
1249/// sequential path.)
1250fn prefill_batched() -> bool {
1251    std::env::var("CMF_PREFILL")
1252        .map(|v| v != "seq")
1253        .unwrap_or(true)
1254}
1255
1256/// Decide the graph NLL route without conflating graph quality with the
1257/// optional native-Metal fused head. A hidden-state graph remains a valid
1258/// quality route on Vulkan/Wgpu; only native Metal requires graph logits.
1259#[inline]
1260fn nll_graph_policy(
1261    unmasked: bool,
1262    prefer_graph: bool,
1263    native_metal: bool,
1264) -> (bool, bool) {
1265    let graph_quality = unmasked && prefer_graph;
1266    let fused_head_quality = graph_quality && native_metal;
1267    (graph_quality, fused_head_quality)
1268}
1269
1270/// Input to the layer-major batched span walk: token ids (embeds itself,
1271/// full-stack and coordinator prefill) or ready boundary hiddens (the
1272/// network worker's side of a split).
1273#[derive(Clone, Copy)]
1274enum PrefillIn<'a> {
1275    Ids(&'a [u32]),
1276    Hidden(&'a [f32]),
1277}
1278
1279/// The batched prefill walks `weights.layers`. Architectures that load
1280/// their own stack (gemma-3n's AltUp replicas, DeepSeek-V4's hyper-
1281/// connections) leave that empty and must go position by position — asking
1282/// otherwise indexes an empty vector, which is a panic rather than a
1283/// fallback. Every call site goes through here so the next such
1284/// architecture is one line, not four.
1285impl Pipeline {
1286    fn can_prefill_batched(&self) -> bool {
1287        #[cfg(test)]
1288        let force_serial = self.nll_test_force_serial;
1289        #[cfg(not(test))]
1290        let force_serial = false;
1291        prefill_batched() && !force_serial && !self.weights.layers.is_empty()
1292    }
1293
1294    /// The backend's automatic capacity split for a mapped transformer.
1295    /// Kept as a method so prefill and decode use the exact same boundary.
1296    fn automatic_gpu_prefix(&self) -> Option<usize> {
1297        let (model, _, _, _) = self.weights.embed_tokens.graph_weight()?;
1298        crate::gpu::automatic_layer_prefix(&model, self.num_layers, self.physical_layers)
1299    }
1300
1301    /// Positions per batched pass of the layer-stack prefill for THIS
1302    /// model on THIS backend (see [`prefill_chunk_rule`]). Pub: the network
1303    /// split must chunk exactly like the local path to reproduce it.
1304    pub fn prefill_chunk(&self) -> usize {
1305        let env = env_prefill_chunk();
1306        if env.is_some() || ChunkHost::here() != ChunkHost::Other {
1307            return prefill_chunk_rule(env, ChunkHost::here(), false);
1308        }
1309        prefill_chunk_rule(None, ChunkHost::Other, self.chunk_stack_facts().dense_on_discrete())
1310    }
1311
1312    fn chunk_stack_facts(&self) -> ChunkStackFacts {
1313        let plain_dense = !self.weights.layers.is_empty()
1314            && self.g3n.is_none()
1315            && self.dsv4.is_none()
1316            && self.dsv41.is_none()
1317            && self.qwen4_exp.is_none()
1318            && self.weights.layers.iter().all(|lw| {
1319                matches!(lw.attn, AttnKind::Full { .. }) && matches!(lw.ffn, FfnKind::Dense(_))
1320            });
1321        let gpu_on = crate::gpu::enabled();
1322        ChunkStackFacts {
1323            plain_dense,
1324            discrete: gpu_on && crate::gpu::discrete(),
1325            gpu_on,
1326            // Only asked when the rest already qualifies: it opens the
1327            // backend's capacity plan.
1328            capacity_split: std::env::var_os("CMF_GPU_LAYERS").is_some()
1329                || (plain_dense && gpu_on && self.automatic_gpu_prefix().is_some()),
1330            multi_gpu: self.gpu_plan.is_some(),
1331            o1: self.o1_active(),
1332        }
1333    }
1334}
1335
1336/// Prefill chunk (positions per batched pass), model-agnostic form. On
1337/// macOS the AMX GEMM path wants tall panels — M=48 starves the matrix
1338/// units (ggml uses ubatch 512); elsewhere the historical 48 stays.
1339/// CMF_PREFILL_CHUNK overrides. The architectures with their own stacks
1340/// (DeepSeek-V4/V4.1) chunk with this; the layer-stack prefill asks
1341/// [`Pipeline::prefill_chunk`], which also knows the model and the card.
1342/// A different chunk is a different (equally valid) generation: panel
1343/// width reorders float accumulation.
1344pub fn prefill_chunk() -> usize {
1345    prefill_chunk_rule(env_prefill_chunk(), ChunkHost::here(), false)
1346}
1347
1348fn env_prefill_chunk() -> Option<usize> {
1349    std::env::var("CMF_PREFILL_CHUNK")
1350        .ok()
1351        .and_then(|v| v.parse::<usize>().ok())
1352}
1353
1354/// The host classes the chunk width distinguishes.
1355#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1356enum ChunkHost {
1357    Macos,
1358    /// Linux/Android aarch64 (phones, SBCs).
1359    Aarch64,
1360    /// Everything else: x86-64 Linux/Windows, CPU or Vulkan/DX12.
1361    Other,
1362}
1363
1364impl ChunkHost {
1365    fn here() -> Self {
1366        if cfg!(target_os = "macos") {
1367            ChunkHost::Macos
1368        } else if cfg!(target_arch = "aarch64") {
1369            ChunkHost::Aarch64
1370        } else {
1371            ChunkHost::Other
1372        }
1373    }
1374}
1375
1376/// Chunk for a plain dense stack whose every layer lives on a discrete
1377/// card. On x86 the layer-stack prefill is host-driven: each GEMM and the
1378/// chunk attention (which re-uploads the whole KV prefix per layer) is a
1379/// separate submit + readback, so 48 positions a pass left the card idle
1380/// between them. Measured in-process on an RTX 3090 (Vulkan), 2048-token
1381/// prompt — see CHANGELOG 0.7.6 for the table.
1382const DISCRETE_DENSE_PREFILL_CHUNK: usize = 512;
1383
1384/// The chunk-width rule. `dense_on_discrete` is true only for a plain
1385/// dense transformer (full attention, dense FFN, no special stack) that
1386/// is entirely resident on one discrete card — the one case measured
1387/// here. GDN hybrids, MoE, DeepSeek stacks, capacity-split and CPU-only
1388/// runs keep the width they were tuned with.
1389fn prefill_chunk_rule(env: Option<usize>, host: ChunkHost, dense_on_discrete: bool) -> usize {
1390    if let Some(n) = env {
1391        return n.max(1);
1392    }
1393    match host {
1394        ChunkHost::Macos => 512,
1395        // Mobile: big enough to feed the batched attend (gate b ≥ 32)
1396        // and the blocked SDOT GEMM without the memory of 512.
1397        ChunkHost::Aarch64 => 256,
1398        ChunkHost::Other if dense_on_discrete => DISCRETE_DENSE_PREFILL_CHUNK,
1399        ChunkHost::Other => 48,
1400    }
1401}
1402
1403/// What the chunk rule needs to know about a loaded stack.
1404#[derive(Clone, Copy, Debug, Default)]
1405struct ChunkStackFacts {
1406    /// Every layer is `AttnKind::Full` + `FfnKind::Dense`, and no
1407    /// architecture-owned stack (g3n, DeepSeek-V4/V4.1, qwen4-exp) is set.
1408    plain_dense: bool,
1409    /// The active GPU backend is a discrete card.
1410    discrete: bool,
1411    /// The backend is up and not paused.
1412    gpu_on: bool,
1413    /// A capacity-derived device prefix: some layers run on the host.
1414    capacity_split: bool,
1415    /// An in-process multi-GPU plan is set.
1416    multi_gpu: bool,
1417    /// O(1) layers (their Q trace is recorded by the prefill).
1418    o1: bool,
1419}
1420
1421impl ChunkStackFacts {
1422    fn dense_on_discrete(self) -> bool {
1423        self.plain_dense
1424            && self.discrete
1425            && self.gpu_on
1426            && !self.capacity_split
1427            && !self.multi_gpu
1428            && !self.o1
1429    }
1430}
1431
1432/// Number of prompt rows that have a real teacher-forced next-token pair in a
1433/// prefill span.  The final prompt row has no successor token, so it must not
1434/// be handed to the MTP warm-up.  Keeping this arithmetic in one helper makes
1435/// the full-chunk and tail-chunk boundaries explicit for both the graph and
1436/// CPU implementations.
1437#[inline]
1438fn mtp_prefill_pair_count(start: usize, end: usize, input_len: usize) -> usize {
1439    if end <= start || start >= input_len {
1440        return 0;
1441    }
1442    let rows = (end.min(input_len) - start).min(input_len - start);
1443    if end < input_len {
1444        rows
1445    } else {
1446        rows.saturating_sub(1)
1447    }
1448}
1449
1450/// Callback for streaming tokens. Return `false` to cancel.
1451pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
1452
1453/// One layer's cache ownership at a cross-turn KV reuse boundary.
1454#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1455pub(crate) struct ReuseLayer {
1456    /// Exact-attention layer (rows in `LayerKvCache`); otherwise a
1457    /// recurrent / latent mixer whose state cannot be rewound.
1458    pub full: bool,
1459    /// Rows the host owner cache holds.
1460    pub host_rows: usize,
1461    /// Rows the wgpu token graph's device mirror holds (None: no mirror).
1462    pub device_rows: Option<usize>,
1463    /// A recurrent state lives on the device (advanced past the host copy).
1464    pub device_state: bool,
1465}
1466
1467/// What a reused turn must do before its tail prefill runs on the HOST.
1468#[derive(Debug, Clone, PartialEq, Eq)]
1469pub(crate) enum ReusePlan {
1470    /// Host caches already hold exactly the reused prefix.
1471    Ready,
1472    /// Copy device mirror rows `[from..to)` into the host cache of each
1473    /// listed layer (the rows decode wrote on the device only).
1474    Pull(Vec<(usize, usize, usize)>),
1475    /// The prefix cannot be continued on the host exactly: start fresh.
1476    Fresh,
1477}
1478
1479/// The wgpu whole-token graph decodes into a DEVICE K/V mirror and never
1480/// writes those rows back to the host cache, while the chunked prefill of a
1481/// pure-attention model reads (and appends to) the host cache. A reused turn
1482/// therefore found its host cache ending at the previous PROMPT, not at the
1483/// previous answer: the tail prefill attended without the model's own
1484/// answer and appended its rows at the wrong index (MiniCPM5 on Vulkan
1485/// repeated its tool call instead of reading the tool result). Every layer
1486/// must hold exactly `reuse_from` host rows before the host continues; rows
1487/// that exist only on the device are pulled back, anything else is fresh.
1488pub(crate) fn kv_reuse_plan(reuse_from: usize, layers: &[ReuseLayer]) -> ReusePlan {
1489    let mut pulls = Vec::new();
1490    for (li, l) in layers.iter().enumerate() {
1491        if !l.full {
1492            if l.device_state {
1493                return ReusePlan::Fresh;
1494            }
1495            continue;
1496        }
1497        if l.host_rows == reuse_from {
1498            continue;
1499        }
1500        if l.host_rows < reuse_from && l.device_rows.is_some_and(|d| d >= reuse_from) {
1501            pulls.push((li, l.host_rows, reuse_from));
1502            continue;
1503        }
1504        return ReusePlan::Fresh;
1505    }
1506    if pulls.is_empty() {
1507        ReusePlan::Ready
1508    } else {
1509        ReusePlan::Pull(pulls)
1510    }
1511}
1512
1513impl Pipeline {
1514    /// Clear all per-sequence state, including backend device mirrors.
1515    ///
1516    /// The host KV/history buffers are only half of the request lifecycle on
1517    /// wgpu: GDN/O(1) state and cached graph bind groups are keyed by the
1518    /// pipeline id and otherwise survive a pooled request.  Keep every fresh
1519    /// sequence entry point on this one reset path so a new request cannot
1520    /// inherit the prior request's device state.
1521    fn clear_sequence_state(&mut self) {
1522        // a replay still writing the GDN owners must land before they are
1523        // cleared or reallocated (the device holds raw pointers to them)
1524        #[cfg(target_os = "macos")]
1525        let _ = crate::gpu_metal::wait_replay();
1526        self.kv_cache.clear();
1527        // Both reuse keys (the legacy `kv_history` and the bounded
1528        // `kv_prefix`) describe the state being dropped here.
1529        self.clear_history();
1530        self.graph_logits = None;
1531        if let Some(b) = &mut self.dsv41 {
1532            b.3.clear();
1533        }
1534        crate::gpu::graph_kv_reset(self.graph_kv_id);
1535        // MTP is detached from `self` for the duration of generation, so its
1536        // device mirror is not covered by the trunk reset above.  Reset the
1537        // derived id as well: a failed/aborted warm-up must never leave a
1538        // mirror that a later request can mistake for a current MTP cache.
1539        crate::gpu::graph_kv_reset(self.mtp_kv_id());
1540    }
1541
1542    /// Make the host caches own exactly the reused prefix `[0..reuse_from)`
1543    /// before a reused turn's tail prefill runs on the host (see
1544    /// [`kv_reuse_plan`]). Returns false when the prefix cannot be continued
1545    /// exactly — the caller then starts a fresh sequence. A model whose
1546    /// prefill runs through the token graph keeps its device state as the
1547    /// authority and is left untouched.
1548    fn prepare_kv_reuse(&mut self, reuse_from: usize) -> bool {
1549        if self.graph_prefill_preferred() {
1550            return true;
1551        }
1552        let kv_id = self.graph_kv_id;
1553        let layers: Vec<ReuseLayer> = (0..self.num_layers)
1554            .map(|li| {
1555                let full = matches!(
1556                    self.weights.layers[self.phys_layer(li)].attn,
1557                    AttnKind::Full { .. }
1558                );
1559                ReuseLayer {
1560                    full,
1561                    host_rows: self.kv_cache.layers[li].seq_len,
1562                    device_rows: crate::gpu::graph_kv_stored(kv_id, li),
1563                    device_state: crate::gpu::graph_state_resident(kv_id, li),
1564                }
1565            })
1566            .collect();
1567        // No wgpu device state at all (CPU, Metal — whose graph appends every
1568        // decoded row to the owner cache itself): the host is the owner and
1569        // the extension check already proved the prefix.
1570        if layers
1571            .iter()
1572            .all(|l| l.device_rows.is_none() && !l.device_state)
1573        {
1574            return true;
1575        }
1576        let plan = kv_reuse_plan(reuse_from, &layers);
1577        let (what, rows, n) = match &plan {
1578            ReusePlan::Ready => ("host ready", 0, 0),
1579            ReusePlan::Fresh => ("fresh", 0, 0),
1580            ReusePlan::Pull(p) => (
1581                "pull",
1582                p.iter().map(|&(_, a, b)| b - a).max().unwrap_or(0),
1583                p.len(),
1584            ),
1585        };
1586        let t0 = std::time::Instant::now();
1587        let ok = self.apply_kv_reuse_plan(reuse_from, plan, &layers);
1588        if std::env::var("CMF_PREFILL_PROF").is_ok() {
1589            eprintln!(
1590                "kv-reuse: {what}{}: {rows} device row(s) × {n} layer(s) to the host in {:.2} ms",
1591                if ok { "" } else { " (failed → fresh)" },
1592                t0.elapsed().as_secs_f64() * 1e3
1593            );
1594        }
1595        ok
1596    }
1597
1598    fn apply_kv_reuse_plan(
1599        &mut self,
1600        reuse_from: usize,
1601        plan: ReusePlan,
1602        layers: &[ReuseLayer],
1603    ) -> bool {
1604        let kv_id = self.graph_kv_id;
1605        match plan {
1606            ReusePlan::Fresh => return false,
1607            ReusePlan::Ready => {}
1608            ReusePlan::Pull(pulls) => {
1609                // Mirrors of one uniform geometry: one batched read serves
1610                // every layer. Per-layer geometry (MiMo-V2: 4/8 KV heads,
1611                // narrow V, sliding rings) is read layer by layer in the
1612                // host layout instead.
1613                let (nkv, hd) = {
1614                    let c = &self.kv_cache.layers[pulls[0].0];
1615                    (c.num_kv_heads, c.head_dim)
1616                };
1617                let uniform = pulls.iter().all(|&(li, _, _)| {
1618                    let c = &self.kv_cache.layers[li];
1619                    (c.num_kv_heads, c.head_dim) == (nkv, hd)
1620                });
1621                let batched = if uniform {
1622                    crate::gpu::graph_kv_read_rows(kv_id, &pulls, nkv, hd)
1623                } else {
1624                    None
1625                };
1626                let rows: Vec<(Vec<f32>, Vec<f32>)> = match batched {
1627                    Some(rows) => rows,
1628                    None => {
1629                        let mut rows = Vec::with_capacity(pulls.len());
1630                        for &(li, from, to) in &pulls {
1631                            let (lnkv, lhd) = {
1632                                let c = &self.kv_cache.layers[li];
1633                                (c.num_kv_heads, c.head_dim)
1634                            };
1635                            let Some((k, v, first_valid)) =
1636                                crate::gpu::graph_kv_pull_host(kv_id, li, from, to, lnkv, lhd)
1637                            else {
1638                                return false;
1639                            };
1640                            // The host continues at `to`: a sliding layer
1641                            // reads back only its last window, a full one
1642                            // every row it lacks.
1643                            let need_from = match self.layer_window(li) {
1644                                Some(w) => from.max((to + 1).saturating_sub(w)),
1645                                None => from,
1646                            };
1647                            if first_valid > need_from {
1648                                return false;
1649                            }
1650                            rows.push((k, v));
1651                        }
1652                        rows
1653                    }
1654                };
1655                for ((li, from, to), (k, v)) in pulls.into_iter().zip(rows) {
1656                    let cache = &mut self.kv_cache.layers[li];
1657                    let row = cache.num_kv_heads * cache.head_dim;
1658                    for p in 0..to - from {
1659                        cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
1660                    }
1661                    if cache.seq_len != to {
1662                        return false;
1663                    }
1664                }
1665            }
1666        }
1667        // A mirror past the prefix (a greedy burst that ran beyond the stop)
1668        // holds rows of the OLD continuation: rewind it so the next graph
1669        // token re-syncs those positions from the host.
1670        for (li, l) in layers.iter().enumerate() {
1671            if l.full
1672                && l.device_rows.is_some_and(|d| d > reuse_from)
1673                && !crate::gpu::graph_kv_set_stored(kv_id, li, reuse_from)
1674            {
1675                return false;
1676            }
1677        }
1678        true
1679    }
1680
1681    /// Finish a generation lifecycle after the MTP/router owners were
1682    /// detached.  Every terminal path must put those owners back before the
1683    /// pooled pipeline can serve another request.  Graph side channels and
1684    /// device mirrors are cleared on errors and cancellations; a successful
1685    /// generation keeps its decode-ready host cache for KV reuse.
1686    fn finish_generation(
1687        &mut self,
1688        mtp: &mut Option<MtpModule>,
1689        router: &mut Option<crate::swarm::DynRouter>,
1690        clear_sequence: bool,
1691    ) {
1692        // A dynamic route may have switched the overlay before the terminal
1693        // path. Restore the backbone while the detached router is still
1694        // available, because set_active_skill also owns the overlay reset.
1695        if router.is_some() {
1696            let _ = self.set_active_skill(None);
1697        }
1698        // The last speculative round's replay may still be in flight on
1699        // the second queue: whoever reads the host cache after generate()
1700        // returns (session export, the network split's KV wire, a KV
1701        // reuse) must see the final states.
1702        // A replay that failed leaves the GDN owners half-written: fail
1703        // closed and drop the sequence instead of handing the cache on.
1704        #[cfg(target_os = "macos")]
1705        let clear_sequence = clear_sequence || !crate::gpu_metal::wait_replay();
1706        if clear_sequence {
1707            self.clear_sequence_state();
1708            if let Some(m) = mtp.as_mut() {
1709                // The MTP owner is detached while generation runs, so the
1710                // trunk reset above cannot clear its host cache.  Drop its
1711                // partial rows before reattaching it to the pooled pipeline;
1712                // the next request must start from the same empty anchor on
1713                // CPU and on the device mirror.
1714                m.kv.clear();
1715            }
1716            if let Some(m) = self.mtp.as_mut() {
1717                // A non-speculative request leaves the configured MTP owner
1718                // attached.  Clear that dormant cache too when a shared
1719                // generation failure/cancellation resets the sequence.
1720                m.kv.clear();
1721            }
1722        }
1723        self.graph_want_logits = false;
1724        self.graph_head_required = false;
1725        self.graph_logits = None;
1726        self.graph_failed
1727            .store(false, std::sync::atomic::Ordering::Relaxed);
1728        self.cancel
1729            .store(false, std::sync::atomic::Ordering::Relaxed);
1730        self.dyn_router = router.take().or(self.dyn_router.take());
1731        self.mtp = mtp.take().or(self.mtp.take());
1732        self.mtp_graph_mode = None;
1733        self.spec_forced = None;
1734    }
1735
1736    /// Consume a graph failure reported by a forward that returns only a
1737    /// hidden vector.  `forward_ids` is a public Result API, so it must not
1738    /// turn the graph's zero hidden sentinel into a valid lm_head result.
1739    fn check_forward_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1740        if self
1741            .graph_failed
1742            .swap(false, std::sync::atomic::Ordering::Relaxed)
1743        {
1744            self.cancel
1745                .store(false, std::sync::atomic::Ordering::Relaxed);
1746            self.clear_sequence_state();
1747            self.graph_logits = None;
1748            self.graph_want_logits = false;
1749            self.graph_head_required = false;
1750            return Err(format!("GPU graph failed during {phase} at position {pos}"));
1751        }
1752        Ok(())
1753    }
1754
1755    #[cfg(target_os = "macos")]
1756    fn fail_metal_graph(&mut self, reason: &str) {
1757        crate::pipeline::METAL_GRAPH_ERRORS
1758            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1759        self.clear_sequence_state();
1760        self.graph_logits = None;
1761        self.graph_failed
1762            .store(true, std::sync::atomic::Ordering::Relaxed);
1763        self.cancel
1764            .store(true, std::sync::atomic::Ordering::Relaxed);
1765        tracing::error!("native Metal TokenGraph failed closed: {reason}");
1766    }
1767
1768    /// Start an NLL/PPL request with all graph side channels in a known
1769    /// state.  A graph failure also raises the cooperative cancel bit; it is
1770    /// consumed here and that graph-induced bit is cleared so an independent
1771    /// request can be reused.  A caller-owned cancellation remains intact.
1772    fn nll_begin(&mut self) -> Result<(), String> {
1773        if self
1774            .graph_failed
1775            .swap(false, std::sync::atomic::Ordering::Relaxed)
1776        {
1777            self.cancel
1778                .store(false, std::sync::atomic::Ordering::Relaxed);
1779            self.clear_sequence_state();
1780            self.graph_logits = None;
1781            self.graph_want_logits = false;
1782            self.graph_head_required = false;
1783            return Err("GPU graph failed before NLL scoring".to_string());
1784        }
1785        self.clear_sequence_state();
1786        self.graph_logits = None;
1787        self.graph_want_logits = false;
1788        self.graph_head_required = false;
1789        Ok(())
1790    }
1791
1792    /// End an NLL/PPL request, including the side channels that are not part
1793    /// of the host KV cache.  This is intentionally explicit instead of
1794    /// relying on a tuple/sentinel return: callers must see every failure.
1795    fn nll_end(&mut self) {
1796        self.clear_sequence_state();
1797        self.graph_logits = None;
1798        self.graph_want_logits = false;
1799        self.graph_head_required = false;
1800        self.graph_failed
1801            .store(false, std::sync::atomic::Ordering::Relaxed);
1802    }
1803
1804    /// Check the graph failure channel at a scoring boundary and leave the
1805    /// pipeline reusable when the device path failed.
1806    fn nll_check_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1807        #[cfg(test)]
1808        if self.nll_test_fail_at == Some(pos) {
1809            self.nll_test_fail_at = None;
1810            self.graph_failed
1811                .store(true, std::sync::atomic::Ordering::Relaxed);
1812            self.cancel
1813                .store(true, std::sync::atomic::Ordering::Relaxed);
1814        }
1815        if self
1816            .graph_failed
1817            .swap(false, std::sync::atomic::Ordering::Relaxed)
1818        {
1819            self.cancel
1820                .store(false, std::sync::atomic::Ordering::Relaxed);
1821            self.clear_sequence_state();
1822            self.graph_logits = None;
1823            self.graph_want_logits = false;
1824            return Err(format!(
1825                "GPU graph failed during NLL {phase} at position {pos}"
1826            ));
1827        }
1828        Ok(())
1829    }
1830
1831    /// Map a virtual layer index to its physical weight index.
1832    /// Looped Transformer (Nanbeige 4.2): 22 physical layers × 2 loops = 44 virtual;
1833    /// virtual layer 23 maps back to physical layer 1 (23 % 22 = 1).
1834    #[inline]
1835    pub fn phys_layer(&self, virtual_idx: usize) -> usize {
1836        virtual_idx % self.physical_layers
1837    }
1838
1839    /// True when `virtual_idx` is the last layer of a loop iteration
1840    /// (used for loop_final_norm insertion).
1841    #[inline]
1842    pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
1843        self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
1844    }
1845
1846    /// Build a pipeline from parts (used by the loader and tests).
1847    #[allow(clippy::too_many_arguments)]
1848
1849    /// Whole-block q1 token graph on the GPU (macOS/Metal): the run of
1850    /// consecutive q1 layers — GDN *and* full attention — starting at
1851    /// `start` executes as few command buffers as the CPU truly needs.
1852    /// Hidden stays device-resident across every layer; the only syncs
1853    /// are before each CPU attend (it needs q/k/v and owns the KV
1854    /// cache) and the final hidden readback. Recurrent states
1855    /// round-trip through shared memory (the CPU stays their owner, so
1856    /// every other path remains coherent). Returns the first layer
1857    /// index NOT covered (== `start` → refused, caller falls through
1858    /// to the per-layer CPU path).
1859    /// Should prefill run position-by-position through the GPU token
1860    /// graph instead of the batched CPU chunk-GEMM? True for q1 GDN
1861    /// hybrids on native Metal: their chunk prefill is walled by the
1862    /// sequential scalar recurrence, so the graph's decode rate wins.
1863    /// NOT for Looped Transformers, despite the per-chunk loop_final_norm
1864    /// sync: the chunk-GEMM amortizes each weight over the whole chunk,
1865    /// which the per-position graph cannot (Nanbeige 4.2 on M4, 512-token
1866    /// prompt: 85 tok/s chunked vs 14 through the graph).
1867    #[cfg(target_os = "macos")]
1868    fn graph_prefill_preferred(&self) -> bool {
1869        let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
1870        if !crate::gpu::enabled_here()
1871            || !graph_force
1872            || std::env::var("CMF_GPU_BLOCK")
1873                .map(|v| v == "0")
1874                .unwrap_or(false)
1875            // CMF_PREFILL_GRAPH=0: the chunked prefill (GEMM projections,
1876            // CPU recurrence) instead of the per-position token graph.
1877            || std::env::var("CMF_PREFILL_GRAPH").as_deref() == Ok("0")
1878        {
1879            return false;
1880        }
1881        self.weights
1882            .layers
1883            .iter()
1884            .any(|lw| {
1885                matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.metal_graph_parts().is_some())
1886            })
1887    }
1888
1889    /// Prompt ingest through the batched wgpu graph in device-prefix mode:
1890    /// a MoE stack that does not fit the card runs each chunk's leading
1891    /// layers on the device (experts resident) and the rest on the host's
1892    /// batched walk. On by default for models with per-layer attention
1893    /// geometry (MiMo-V2 — its measured default); `CMF_BATCH_PREFIX=1`
1894    /// opts any other MoE model in, `=0` keeps the chunked host prefill.
1895    #[cfg(not(target_os = "macos"))]
1896    fn batch_prefix_prefill(&self) -> bool {
1897        let forced = match std::env::var("CMF_BATCH_PREFIX").as_deref() {
1898            Ok("0") => return false,
1899            Ok("1") => true,
1900            _ => false,
1901        };
1902        crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
1903            && crate::gpu::enabled_here()
1904            && !self.graph_refused()
1905            && (forced || self.graph_attn_decline_reason().is_some())
1906            && self.wgpu_graph_attn_decline().is_none()
1907            && self.attn_softcap == 0.0
1908            && self
1909                .weights
1910                .layers
1911                .iter()
1912                .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
1913            && self.automatic_gpu_prefix().is_some()
1914    }
1915
1916    #[cfg(not(target_os = "macos"))]
1917    fn graph_prefill_preferred(&self) -> bool {
1918        // Discrete-GPU wgpu whole-token graph: GDN layers carry recurrent state
1919        // (conv ring + delta-rule S) resident on the GPU. A batched CPU prefill
1920        // builds that state on the CPU only, leaving the GPU buffers zeroed at
1921        // decode → garbage. Route GDN-hybrid prefill through the graph one
1922        // position at a time so the resident state is seeded exactly as decode
1923        // will read it. Pure-attention models keep the batched CPU prefill (its
1924        // KV mirror re-syncs from the CPU cache, so no seeding gap).
1925        let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
1926        if !graph_on || !crate::gpu::enabled_here() {
1927            return false;
1928        }
1929        // Embryo's phase state is device-owned by the resident graph during
1930        // prefill; the batched CPU path would leave decode seeing a zeroed
1931        // device recurrence.  Route the prompt position-by-position too.
1932        if self.embryo_resident_eligible() {
1933            // The resident graph owns the phase state and anchor KV on the
1934            // device.  A prefill-only graph would leave decode on the host
1935            // with no way to import that state, so keep the whole sequence
1936            // on one owner (or use the ordinary CPU prefill/decode pair).
1937            return crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
1938        }
1939        // The descriptor-aware Prism graph now carries both the FWHT/affine
1940        // transforms and resident GDN state, so it is also the exact prefill
1941        // path for this model.  Keeping it here (rather than falling through
1942        // to the CPU chunk walk) is required for a long prompt to seed the
1943        // same device state that decode consumes.
1944        // O(1) needs the CPU prefill: the q-trace that seals the Nyström
1945        // skeleton is recorded there and nowhere else. The GDN half of
1946        // the hybrid loses nothing — the graph's first decode creates
1947        // its (ring, S) entries seeded from `cpu_state`, the same
1948        // handoff every graph run relies on when the entry is fresh.
1949        // Without this line the two designs collide on hybrids and o1
1950        // never becomes graph-portable: prefill through the graph
1951        // records no trace, so views stay None forever.
1952        if self.o1_active() {
1953            return false;
1954        }
1955        // A model the wgpu graphs decline outright would walk its prompt
1956        // one position at a time through a graph that never runs: the
1957        // batched CPU chunk prefill is the right ingest for it.
1958        if self.wgpu_graph_attn_decline().is_some() {
1959            return false;
1960        }
1961        if self
1962            .weights
1963            .layers
1964            .iter()
1965            .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
1966        {
1967            return true;
1968        }
1969        // MoE models too: the chunked CPU prefill runs every expert on the
1970        // host (Hy-MT2-30B-A3B on a Xeon: 8 tok/s of ingest against 53 of
1971        // graph decode), while the token graph — and the batched graph under
1972        // CMF_BATCH_K — keep the experts resident. Full attention in the
1973        // graph writes the KV mirror that decode reads, exactly as it does
1974        // for the hybrids' attention layers. Only when the whole stack is
1975        // resident: with a device prefix the per-position walk finishes
1976        // every token on the host, and the chunked prefill (GEMMs on the
1977        // card, the expert loop batched on the host) is the faster ingest
1978        // (the 8 GB ladder point: 7 tok/s chunked against ~1 walked).
1979        self.weights
1980            .layers
1981            .iter()
1982            .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
1983            && self.automatic_gpu_prefix().is_none()
1984    }
1985
1986    #[cfg(target_os = "macos")]
1987    fn q1_graph_gpu(
1988        &mut self,
1989        start: usize,
1990        upto: Option<usize>,
1991        position: usize,
1992        h: &mut [f32],
1993    ) -> usize {
1994        let _mt0 = std::time::Instant::now(); // CMF_METAL_HOSTPROF
1995        use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
1996        let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
1997        if self.attn_softcap > 0.0 // capped scores: no graph kernel — CPU path
1998            || !crate::gpu::enabled_here()
1999            || !graph_force
2000            || std::env::var("CMF_GPU_BLOCK")
2001                .map(|v| v == "0")
2002                .unwrap_or(false)
2003        {
2004            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2005                eprintln!(
2006                    "block-graph: front gate (softcap={} enabled_here={} graph_force={})",
2007                    self.attn_softcap > 0.0,
2008                    crate::gpu::enabled_here(),
2009                    graph_force,
2010                );
2011            }
2012            if self.graph_head_required {
2013                self.fail_metal_graph("native graph front gate refused");
2014            }
2015            return start;
2016        }
2017        // The graph encodes SiLU FFN and full-context attention with an
2018        // explicit model scale. Architectures with sliding windows,
2019        // sandwich norms or non-SiLU FFNs still fall back to the CPU path.
2020        if self.swa.is_some()
2021            || self.global_attn.is_some()
2022            || self.attention_heads_per_layer.is_some()
2023            || self.attn_v_norm
2024            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
2025            || self.graph_attn_decline_reason().is_some()
2026            || self.weights.layers.iter().any(|lw| {
2027                lw.attn_out_norm.is_some()
2028                    || lw.ffn_out_norm.is_some()
2029                    || lw.layer_scale.is_some()
2030                    || matches!(&lw.ffn, FfnKind::Dense(d) if d.act != Act::Silu)
2031            })
2032        {
2033            // The Metal graphs have no per-layer attention geometry (the
2034            // wgpu graphs do): say so once, by name.
2035            if let Some(reason) = self.graph_attn_decline_reason() {
2036                self.note_graph_decline("metal block graph", reason);
2037            }
2038            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2039                eprintln!(
2040                    "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
2041                    self.swa.is_some(),
2042                    self.global_attn.is_some(),
2043                    self.attention_heads_per_layer.is_some(),
2044                    self.attn_v_norm,
2045                    (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
2046                );
2047            }
2048            if self.graph_head_required {
2049                self.fail_metal_graph("native graph architecture gate refused");
2050            }
2051            return start;
2052        }
2053        // Looped Transformer: the graph covers ALL loop iterations;
2054        // encode_loop_norm is inserted on-device at each boundary.
2055        let limit = upto
2056            .map(|u| u + 1)
2057            .unwrap_or(self.num_layers)
2058            .min(self.num_layers);
2059
2060        enum Item<'a> {
2061            Gdn {
2062                run: Vec<GdnGpuLayer<'a>>,
2063                first: usize,
2064            },
2065            Attn {
2066                l: AttnGpuLayer<'a>,
2067                li: usize,
2068                q_norm: Option<&'a [f32]>,
2069                k_norm: Option<&'a [f32]>,
2070                output_gate: bool,
2071                bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
2072                /// Attend on the device too (no sync): F32 KV, no
2073                /// o1/bias, dims inside the kernels' contract.
2074                full_gpu: bool,
2075            },
2076        }
2077
2078        // Device-attend KERNEL contract, shared by every Full layer. The
2079        // hd>128 default-off POLICY is applied after the scan: it was
2080        // measured on dense models, and a MoE plan inverts it — with the
2081        // experts on device each CPU-attend sandwich costs a
2082        // commit+wait, ~30 submits/token (W2 on M4: 14.7 tok/s
2083        // sandwiched vs 27.1 device-attend vs 18.8 pure CPU).
2084        let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
2085        let attend_contract = attend_mode != "0"
2086            && attend_mode != "off"
2087            && self.head_dim % 4 == 0
2088            && self.head_dim <= 256
2089            && self.rotary_dim >= 2
2090            && self.rotary_dim <= self.head_dim
2091            && (self.rotary_dim / 2) % 32 == 0
2092            && self.num_kv_heads > 0
2093            && self.num_heads % self.num_kv_heads == 0;
2094
2095        let mut plan: Vec<Item> = Vec::new();
2096        let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
2097        // Break-reason diagnostics ride the same env as the plan summary.
2098        let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
2099        let mut scan = start;
2100        while scan < limit {
2101            let lw = &self.weights.layers[self.phys_layer(scan)];
2102            let ffn = match &lw.ffn {
2103                FfnKind::Dense(d) if d.segs.is_empty() => {
2104                    let (Some(g), Some(u), Some(dn)) = (
2105                        d.gate_proj.metal_graph_parts(),
2106                        d.up_proj.metal_graph_parts(),
2107                        d.down_proj.metal_graph_parts(),
2108                    ) else {
2109                        if block_diag {
2110                            eprintln!(
2111                                "block-graph: L{scan} FFN trio not graph-mappable — run ends"
2112                            );
2113                        }
2114                        break;
2115                    };
2116                    MetalFfn::Dense {
2117                        gate: g,
2118                        up: u,
2119                        down: dn,
2120                    }
2121                }
2122                FfnKind::Moe(m) => {
2123                    let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
2124                        if block_diag {
2125                            eprintln!(
2126                                "block-graph: L{scan} MoE outside the graph contract — run ends"
2127                            );
2128                        }
2129                        break;
2130                    };
2131                    if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
2132                        model_ref.get_or_insert_with(|| model.clone());
2133                    }
2134                    MetalFfn::Moe(moe)
2135                }
2136                _ => {
2137                    if block_diag {
2138                        eprintln!("block-graph: L{scan} non-graph FFN — run ends");
2139                    }
2140                    break;
2141                }
2142            };
2143            match &lw.attn {
2144                AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
2145                    let parts = (
2146                        w.in_proj_qkv.metal_graph_parts(),
2147                        w.in_proj_z.metal_graph_parts(),
2148                        w.in_proj_a.f32_parts(),
2149                        w.in_proj_b.f32_parts(),
2150                        w.out_proj.metal_graph_parts(),
2151                    );
2152                    let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
2153                        if block_diag {
2154                            eprintln!(
2155                                "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
2156                                w.in_proj_qkv.metal_graph_parts().is_some(),
2157                                w.in_proj_z.metal_graph_parts().is_some(),
2158                                w.in_proj_a.f32_parts().is_some(),
2159                                w.in_proj_b.f32_parts().is_some(),
2160                                w.out_proj.metal_graph_parts().is_some(),
2161                            );
2162                        }
2163                        break;
2164                    };
2165                    if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
2166                        model_ref.get_or_insert_with(|| model.clone());
2167                    }
2168                    let gl = GdnGpuLayer {
2169                        attn_norm: &lw.input_norm,
2170                        post_norm: &lw.post_norm,
2171                        qkv,
2172                        z,
2173                        a,
2174                        b,
2175                        out,
2176                        ffn,
2177                        conv1d: &w.conv1d,
2178                        a_log: &w.a_log,
2179                        dt_bias: &w.dt_bias,
2180                        gnorm: &w.norm,
2181                    };
2182                    match plan.last_mut() {
2183                        Some(Item::Gdn { run, .. }) => run.push(gl),
2184                        _ => plan.push(Item::Gdn {
2185                            run: vec![gl],
2186                            first: scan,
2187                        }),
2188                    }
2189                }
2190                AttnKind::Full {
2191                    wq,
2192                    wk,
2193                    wv,
2194                    wo,
2195                    q_norm,
2196                    k_norm,
2197                    output_gate,
2198                    softplus_gate: None,
2199                    bias,
2200                } if !self.kv_cache.layers[scan].o1_sealed()
2201                    // Sealed o1 stays plannable when the Metal o1 port
2202                    // is on: full_gpu attends through the device state,
2203                    // and any refusal falls to the sandwich, whose CPU
2204                    // core routes sealed layers through the nystrom step.
2205                    || std::env::var("CMF_O1_METAL").as_deref() == Ok("1") =>
2206                {
2207                    let parts = (
2208                        wq.metal_graph_parts(),
2209                        wk.metal_graph_parts(),
2210                        wv.metal_graph_parts(),
2211                        wo.metal_graph_parts(),
2212                    );
2213                    let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
2214                        break;
2215                    };
2216                    if let QTensor::Mapped { model, .. } = wq {
2217                        model_ref.get_or_insert_with(|| model.clone());
2218                    }
2219                    let cache = &self.kv_cache.layers[scan];
2220                    // O(1) layer on Metal: the device attends through the
2221                    // sealed Nystrom state (opt-in while the port proves
2222                    // itself). Unsealed -> sandwich path = the CPU o1 step.
2223                    let o1_metal = cache.o1.is_some()
2224                        && std::env::var("CMF_O1_METAL").as_deref() == Ok("1")
2225                        && cache.o1_views().is_some();
2226                    let full_gpu = attend_contract
2227                        && cache.mode == crate::kv_cache::KvMode::F32
2228                        && (cache.o1.is_none() || o1_metal)
2229                        && bias.is_none()
2230                        && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
2231                        && pk.1 == self.num_kv_heads * self.head_dim
2232                        && pv.1 == self.num_kv_heads * self.head_dim
2233                        && po.2 == self.num_heads * self.head_dim;
2234                    plan.push(Item::Attn {
2235                        l: AttnGpuLayer {
2236                            attn_norm: &lw.input_norm,
2237                            post_norm: &lw.post_norm,
2238                            wq: pq,
2239                            wk: pk,
2240                            wv: pv,
2241                            wo: po,
2242                            ffn,
2243                        },
2244                        li: scan,
2245                        q_norm: q_norm.as_deref(),
2246                        k_norm: k_norm.as_deref(),
2247                        output_gate: *output_gate,
2248                        bias: bias
2249                            .as_ref()
2250                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2251                        full_gpu,
2252                    });
2253                }
2254                _ => break,
2255            }
2256            scan += 1;
2257        }
2258        let Some(model) = model_ref else {
2259            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2260                eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
2261            }
2262            if self.graph_head_required {
2263                self.fail_metal_graph("native graph has no mapped model reference");
2264            }
2265            return start;
2266        };
2267        if plan.is_empty() {
2268            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2269                eprintln!("q1-graph: empty plan at layer {start}");
2270            }
2271            if self.graph_head_required {
2272                self.fail_metal_graph("native graph plan is empty");
2273            }
2274            return start;
2275        }
2276        let has_moe = plan.iter().any(|it| match it {
2277            Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
2278            Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
2279        });
2280        let has_gdn = plan.iter().any(|it| matches!(it, Item::Gdn { .. }));
2281        let dev_attend = attend_contract
2282            && (self.head_dim <= 128
2283                || has_moe
2284                // A GDN hybrid attends on a quarter of its layers: the
2285                // hd>128 caution was measured on pure-dense models where
2286                // gqa_attend dominates, and on Qwen3.8-27B (hd 256, 48
2287                // GDN + 16 attn) the sandwich costs 2x the whole decode
2288                // (1.2 vs 2.21 tok/s measured before the arena fix).
2289                || (self.head_dim <= 256 && has_gdn)
2290                || attend_mode == "force"
2291                || attend_mode == "256");
2292        if !dev_attend {
2293            for it in &mut plan {
2294                if let Item::Attn { li, full_gpu, .. } = it {
2295                    // The hd>128 policy is about gqa_attend; an o1 layer
2296                    // attends through its own kernel set.
2297                    let keep_o1 = self.kv_cache.layers[*li].o1.is_some()
2298                        && std::env::var("CMF_O1_METAL").as_deref() == Ok("1");
2299                    if !keep_o1 {
2300                        *full_gpu = false;
2301                    }
2302                }
2303            }
2304        }
2305        if std::env::var("CMF_GRAPH_DBG").is_ok() {
2306            use std::sync::atomic::{AtomicBool, Ordering};
2307            static SAID: AtomicBool = AtomicBool::new(false);
2308            if !SAID.swap(true, Ordering::Relaxed) {
2309                let fg = plan
2310                    .iter()
2311                    .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
2312                    .count();
2313                let att = plan
2314                    .iter()
2315                    .filter(|it| matches!(it, Item::Attn { .. }))
2316                    .count();
2317                eprintln!(
2318                    "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
2319                    plan.len(),
2320                    self.head_dim,
2321                    self.rotary_dim,
2322                    self.num_kv_heads,
2323                    self.num_heads,
2324                );
2325            }
2326        }
2327        let dims = GraphDims {
2328            hidden: self.hidden_size,
2329            eps: self.rms_eps as f32,
2330            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2331        };
2332        let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
2333            if self.graph_head_required {
2334                self.fail_metal_graph("native TokenGraph allocation refused");
2335            }
2336            return start;
2337        };
2338        let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
2339            nv: cfg.num_v_heads,
2340            nk: cfg.num_k_heads,
2341            dk: cfg.key_head_dim,
2342            dv: cfg.value_head_dim,
2343            kk: cfg.conv_kernel,
2344            hidden: self.hidden_size,
2345            inter: self.intermediate_size,
2346            c_dim: cfg.conv_dim(),
2347            eps: cfg.rms_eps as f32,
2348            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2349        });
2350        // Validate the whole plan BEFORE encoding anything: after the
2351        // first sync a refused layer would leave the token
2352        // half-executed, so truncate to the provably encodable prefix.
2353        let mut valid = 0usize;
2354        let mut end = start;
2355        crate::gpu::stageprof(1, _mt0.elapsed()); // конец планирования
2356        if std::env::var("CMF_PLAN_DUMP").is_ok() {
2357            static ONCE: std::sync::Once = std::sync::Once::new();
2358            ONCE.call_once(|| {
2359                for it in &plan {
2360                    match it {
2361                        Item::Gdn { first, run } => {
2362                            eprintln!("plan: Gdn first={first} len={}", run.len())
2363                        }
2364                        Item::Attn { li, full_gpu, .. } => {
2365                            eprintln!("plan: Attn li={li} full_gpu={full_gpu}")
2366                        }
2367                    }
2368                }
2369            });
2370        }
2371        for item in &plan {
2372            let ok = match item {
2373                Item::Gdn { run, .. } => gcfg
2374                    .as_ref()
2375                    .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
2376                    .unwrap_or(false),
2377                Item::Attn { l, .. } => graph.attn_ok(l),
2378            };
2379            if !ok {
2380                if block_diag {
2381                    eprintln!(
2382                        "block-graph: plan item {} ({}) failed graph preflight",
2383                        valid,
2384                        match item {
2385                            Item::Gdn { run, first } => format!("GDN run L{first}+{}", run.len()),
2386                            Item::Attn { li, .. } => format!("Attn L{li}"),
2387                        }
2388                    );
2389                }
2390                break;
2391            }
2392            valid += 1;
2393            end += match item {
2394                Item::Gdn { run, .. } => run.len(),
2395                Item::Attn { .. } => 1,
2396            };
2397        }
2398        plan.truncate(valid);
2399        if plan.is_empty() {
2400            if self.graph_head_required {
2401                self.fail_metal_graph("native graph preflight produced no valid items");
2402            }
2403            return start;
2404        }
2405
2406        if self.graph_head_required && (upto.is_some() || end != self.num_layers) {
2407            self.fail_metal_graph("fused-head NLL requires a complete 64-layer graph");
2408            return start;
2409        }
2410
2411        // Plain dense decode (every item a device-attended full-attention
2412        // layer with a dense FFN, no O(1) state): the only plan shape the
2413        // masked-nibble q4tp matvec and the concurrent layer encoder were
2414        // measured on (MiniCPM5-2B, Qwen3-0.6B on the M4). Hybrids, MoE and
2415        // o1 layers keep the historical serial path bit for bit.
2416        // Every projection must be ONE dispatch (q1t adds an overlay pass,
2417        // Prism q2tp a transform pass — dependent pairs a concurrent
2418        // encoder would race).
2419        let one_pass = |t: (usize, usize, usize)| {
2420            use cortiq_core::TensorDtype as D;
2421            matches!(
2422                model.tensors[t.0].dtype,
2423                D::Q4TiledP | D::Q4Tiled | D::Q4Block | D::Q8Row | D::Q8_2f | D::Q1
2424            )
2425        };
2426        let dense_fast = plan.iter().all(|it| match it {
2427            Item::Attn {
2428                l, li, full_gpu, ..
2429            } => {
2430                *full_gpu
2431                    && self.kv_cache.layers[*li].o1.is_none()
2432                    && [l.wq, l.wk, l.wv, l.wo].into_iter().all(one_pass)
2433                    && match l.ffn {
2434                        MetalFfn::Dense { gate, up, down } => {
2435                            one_pass(gate) && one_pass(up) && one_pass(down)
2436                        }
2437                        _ => false,
2438                    }
2439            }
2440            Item::Gdn { .. } => false,
2441        });
2442        let ab = crate::gpu_metal::dense_ab_arm().filter(|_| dense_fast);
2443        let _mv_fast = match ab {
2444            Some((bits, _)) => {
2445                graph.set_dense_concurrent_raw(bits & crate::gpu_metal::DENSE_CONC != 0);
2446                crate::gpu_metal::MvFastGuard::set_raw(bits)
2447            }
2448            None => {
2449                graph.set_dense_concurrent(dense_fast);
2450                crate::gpu_metal::MvFastGuard::set_bits(if dense_fast {
2451                    crate::gpu_metal::DENSE_MV | crate::gpu_metal::DENSE_FUSE
2452                } else {
2453                    0
2454                })
2455            }
2456        };
2457
2458        let inv_freq = self.inv_freq.clone();
2459        let pool = self.pool.clone();
2460        let (nh, nkv, hd, hs, rd, eps) = (
2461            self.num_heads,
2462            self.num_kv_heads,
2463            self.head_dim,
2464            self.hidden_size,
2465            self.rotary_dim,
2466            self.rms_eps,
2467        );
2468        let norm_style = self.norm_style;
2469        let gemma = norm_style == cortiq_core::NormStyle::Gemma;
2470        let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
2471        let kv_id = self.graph_kv_id;
2472        // GDN runs whose states await readback after the next sync
2473        // (device-attended layers add no sync, so several may stack).
2474        let mut pending: Vec<(usize, usize)> = Vec::new();
2475        // Device-attended layers: their K/V/imp are pulled from the
2476        // mirror after the final sync.
2477        let mut dev_attn: Vec<usize> = Vec::new();
2478        for item in &plan {
2479            let _xt0 = std::time::Instant::now();
2480            let _xkind: u32 = match item {
2481                Item::Gdn { .. } => 2,
2482                Item::Attn { .. } => 3,
2483            };
2484            // Looped Transformer: insert on-device norm at loop boundaries.
2485            if self.loop_final_norm {
2486                let item_start = match item {
2487                    Item::Gdn { first, .. } => *first,
2488                    Item::Attn { li, .. } => *li,
2489                };
2490                if item_start > start && self.is_loop_end(item_start - 1) {
2491                    graph.encode_loop_norm(&self.weights.final_norm);
2492                }
2493            }
2494            match item {
2495                Item::Gdn { run, first } => {
2496                    for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
2497                        if l.linear_state.len() != want {
2498                            l.linear_state = vec![0f32; want];
2499                        }
2500                    }
2501                    let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
2502                        .iter()
2503                        .map(|l| l.linear_state.as_slice())
2504                        .collect();
2505                    let _ig = std::time::Instant::now();
2506                    if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
2507                        // Unreachable: the plan was validated above.
2508                        tracing::error!("q1 graph: GDN run refused after validation");
2509                        return start;
2510                    }
2511                    // Early commit: the GPU starts the run while the
2512                    // CPU encodes the next layer (nothing to wait on).
2513                    graph.commit_kind = 2;
2514                    graph.commit();
2515                    crate::gpu::stageprof(0, _ig.elapsed());
2516                    pending.push((*first, run.len()));
2517                }
2518                Item::Attn {
2519                    l,
2520                    li,
2521                    q_norm,
2522                    k_norm,
2523                    output_gate,
2524                    bias,
2525                    full_gpu,
2526                } => {
2527                    let _ia = std::time::Instant::now();
2528                    // ── Fully device-resident attention: no sync at all.
2529                    if *full_gpu {
2530                        let cache = &self.kv_cache.layers[*li];
2531                        let o1p = if cache.o1.is_some() {
2532                            match cache.o1_views() {
2533                                Some(views) => Some(crate::gpu::O1AttnParams {
2534                                    views,
2535                                    epoch: self.o1_epoch,
2536                                }),
2537                                // Sealed state gone mid-run: sandwich.
2538                                None => None,
2539                            }
2540                        } else {
2541                            None
2542                        };
2543                        let o1_layer = cache.o1.is_some();
2544                        if o1_layer && o1p.is_none() {
2545                            // fall to the sandwich (CPU o1 step)
2546                        }
2547                        let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
2548                        let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
2549                        let cpu_stored = if o1_layer { 0 } else { cpu_k[0].len() / hd };
2550                        let p = crate::gpu::AttnDeviceParams {
2551                            kv_id,
2552                            layer: *li,
2553                            nh,
2554                            nkv,
2555                            hd,
2556                            rd,
2557                            position,
2558                            scale: self.attn_scale,
2559                            eps: eps as f32,
2560                            gemma,
2561                            late_qk_norm: self.qk_norm_after_rope,
2562                            output_gate: *output_gate,
2563                            q_norm: *q_norm,
2564                            k_norm: *k_norm,
2565                            inv_freq: &inv_freq,
2566                            cpu_k,
2567                            cpu_v,
2568                            cpu_stored,
2569                            o1: o1p,
2570                        };
2571                        let o1_bad = o1_layer && p.o1.is_none();
2572                        if !o1_bad && graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p)
2573                        {
2574                            // o1 layers leave no mirror row to pull.
2575                            if p.o1.is_none() {
2576                                dev_attn.push(*li);
2577                            }
2578                            graph.commit_kind = 3;
2579                            graph.commit();
2580                            // The footer below is skipped by `continue`:
2581                            // account the device-attn item here or its
2582                            // cost hides from the stage profile entirely.
2583                            crate::gpu::stageprof(_xkind, _xt0.elapsed());
2584                            continue;
2585                        }
2586                        // Mirror refused (nothing encoded) → sandwich.
2587                    }
2588                    graph.encode_attn_prefix(l);
2589                    if let Err(err) = graph.sync_checked() {
2590                        self.fail_metal_graph(&err);
2591                        return start;
2592                    }
2593                    if !pending.is_empty() {
2594                        let idxs: Vec<usize> =
2595                            pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2596                        let mut outs: Vec<&mut [f32]> = self
2597                            .kv_cache
2598                            .layers
2599                            .iter_mut()
2600                            .enumerate()
2601                            .filter(|(i, _)| idxs.binary_search(i).is_ok())
2602                            .map(|(_, s)| s.linear_state.as_mut_slice())
2603                            .collect();
2604                        graph.read_states(&mut outs);
2605                    }
2606                    let mut q_raw = attention::take_buf(l.wq.1);
2607                    let mut k = attention::take_buf(l.wk.1);
2608                    let mut v = attention::take_buf(l.wv.1);
2609                    graph.read_qkv(&mut q_raw, &mut k, &mut v);
2610                    let cfg = QwenAttnCfg {
2611                        num_heads: nh,
2612                        num_kv_heads: nkv,
2613                        head_dim: hd,
2614                        hidden_size: hs,
2615                        position,
2616                        inv_freq: &inv_freq,
2617                        rotary_dim: rd,
2618                        scale: self.attn_scale,
2619                        softcap: self.attn_softcap,
2620                        window: None,
2621                        v_norm: false,
2622                        qk_norm_after_rope: self.qk_norm_after_rope,
2623                        q_norm: *q_norm,
2624                        k_norm: *k_norm,
2625                        output_gate: *output_gate,
2626                        softplus_gate: None,
2627                        rope_scale: 1.0,
2628                        bias: *bias,
2629                        rms_eps: eps,
2630                        norm_style,
2631                        pool: pool.as_deref(),
2632                        v_head_dim: hd,
2633                    };
2634                    // CMF_ATTN_ORACLE=1: diff the device attend against
2635                    // this CPU attend on identical inputs (bring-up).
2636                    let oracle = std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")
2637                        || std::env::var("CMF_ATTN_DUMP").is_ok();
2638                    let _ = full_gpu;
2639                    let oracle_in = oracle.then(|| (q_raw.clone(), k.clone(), v.clone()));
2640                    let mut ao = attention::qwen_attention_core(
2641                        q_raw,
2642                        k,
2643                        v,
2644                        &mut self.kv_cache.layers[*li],
2645                        &cfg,
2646                    );
2647                    // CMF_ATTN_DUMP=<dir>: this token's rope'd Q and the layer's whole
2648                    // K/V cache as raw f32 (offline attention-statistics probes:
2649                    // block bounds, mass concentration). Needs CMF_GPU_ATTEND=0.
2650                    if let Ok(dir) = std::env::var("CMF_ATTN_DUMP") {
2651                        if let Some((qr0, k0, v0)) = oracle_in.clone() {
2652                            let (cq, _cg, _ck, _cv) =
2653                                attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2654                            let cache = &self.kv_cache.layers[*li];
2655                            let n = cache.head_keys(0).len() / hd;
2656                            let mut bytes: Vec<u8> = Vec::new();
2657                            for v in [nh as u32, nkv as u32, hd as u32, n as u32, position as u32] {
2658                                bytes.extend_from_slice(&v.to_le_bytes());
2659                            }
2660                            for v in &cq {
2661                                bytes.extend_from_slice(&v.to_le_bytes());
2662                            }
2663                            for g in 0..nkv {
2664                                for v in cache.head_keys(g) {
2665                                    bytes.extend_from_slice(&v.to_le_bytes());
2666                                }
2667                            }
2668                            for g in 0..nkv {
2669                                for v in cache.head_values(g) {
2670                                    bytes.extend_from_slice(&v.to_le_bytes());
2671                                }
2672                            }
2673                            let _ =
2674                                std::fs::write(format!("{dir}/L{li}_pos{position}.bin"), &bytes);
2675                        }
2676                    }
2677                    if let Some((qr0, k0, v0)) =
2678                        oracle_in.filter(|_| std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1"))
2679                    {
2680                        let (cq, _cg, ck, cv) =
2681                            attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2682                        let mut h_now = vec![0f32; hs];
2683                        graph.read_h(&mut h_now);
2684                        let cache = &self.kv_cache.layers[*li];
2685                        let n_after = cache.head_keys(0).len() / hd;
2686                        // A sealed O(1) cache may have no dense current-row
2687                        // entry. The oracle is a debug probe, so let it see
2688                        // zero stored exact rows instead of underflowing.
2689                        let stored = n_after.saturating_sub(1);
2690                        let cpu_k: Vec<&[f32]> = (0..nkv)
2691                            .map(|g| &cache.head_keys(g)[..stored * hd])
2692                            .collect();
2693                        let cpu_v: Vec<&[f32]> = (0..nkv)
2694                            .map(|g| &cache.head_values(g)[..stored * hd])
2695                            .collect();
2696                        let p = crate::gpu::AttnDeviceParams {
2697                            kv_id,
2698                            layer: *li,
2699                            nh,
2700                            nkv,
2701                            hd,
2702                            rd,
2703                            position,
2704                            scale: self.attn_scale,
2705                            eps: eps as f32,
2706                            gemma,
2707                            late_qk_norm: self.qk_norm_after_rope,
2708                            output_gate: *output_gate,
2709                            q_norm: *q_norm,
2710                            k_norm: *k_norm,
2711                            inv_freq: &inv_freq,
2712                            cpu_k,
2713                            cpu_v,
2714                            cpu_stored: stored,
2715                            o1: None,
2716                        };
2717                        if let Some((dq, dk, dv, dao)) = graph.debug_attn_device(l, &p, &h_now) {
2718                            let md = |a: &[f32], b: &[f32]| {
2719                                a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs()))
2720                            };
2721                            let nn = |a: &[f32]| a.iter().map(|x| x * x).sum::<f32>().sqrt();
2722                            eprintln!(
2723                                "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}",
2724                                nn(&cq),
2725                                md(&cq, &dq),
2726                                nn(&ck),
2727                                md(&ck, &dk),
2728                                nn(&cv),
2729                                md(&cv, &dv),
2730                                nn(&ao),
2731                                md(&ao, &dao)
2732                            );
2733                        } else {
2734                            eprintln!("attn-oracle L{li}: device probe declined");
2735                        }
2736                    }
2737                    graph.encode_attn_suffix(l, &ao);
2738                    // Early commit: the GPU starts O+FFN while the CPU
2739                    // encodes the following GDN run / attention prefix.
2740                    graph.commit();
2741                    attention::recycle_buf(&mut ao);
2742                }
2743            }
2744
2745            crate::gpu::stageprof(_xkind, _xt0.elapsed());
2746        }
2747        // Ride the final norm + lm_head in the same command buffer when
2748        // this run reaches the model's end and the caller wants logits:
2749        // the separate per-op lm_head submit (a full round trip) folds
2750        // into the sync that already happens here.
2751        let mut lm_rows = None;
2752        if self.graph_want_logits
2753            && upto.is_none()
2754            && end == self.num_layers
2755            && std::env::var("CMF_GPU_LMHEAD")
2756                .map(|v| v != "0")
2757                .unwrap_or(true)
2758        {
2759            if let Some(lm) = self.weights.lm_head.metal_graph_parts() {
2760                if graph.lm_head_ok(lm) {
2761                    graph.encode_lm_head(&self.weights.final_norm, lm);
2762                    lm_rows = Some(lm.1);
2763                }
2764            }
2765        }
2766        if self.graph_head_required && lm_rows.is_none() {
2767            METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2768            self.fail_metal_graph("fused graph head was requested but not encodable");
2769            return start;
2770        }
2771        let _sy0 = std::time::Instant::now();
2772        if let Err(err) = graph.sync_checked() {
2773            self.fail_metal_graph(&err);
2774            return start;
2775        }
2776        let _rs0 = std::time::Instant::now();
2777        if !pending.is_empty() {
2778            let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2779            let mut outs: Vec<&mut [f32]> = self
2780                .kv_cache
2781                .layers
2782                .iter_mut()
2783                .enumerate()
2784                .filter(|(i, _)| idxs.binary_search(i).is_ok())
2785                .map(|(_, s)| s.linear_state.as_mut_slice())
2786                .collect();
2787            graph.read_states(&mut outs);
2788        }
2789        if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() == Ok("1") {
2790            use std::sync::atomic::{AtomicU64, Ordering};
2791            static SY: AtomicU64 = AtomicU64::new(0);
2792            static RS: AtomicU64 = AtomicU64::new(0);
2793            static N: AtomicU64 = AtomicU64::new(0);
2794            SY.fetch_add((_rs0 - _sy0).as_nanos() as u64, Ordering::Relaxed);
2795            RS.fetch_add(_rs0.elapsed().as_nanos() as u64, Ordering::Relaxed);
2796            let n = N.fetch_add(1, Ordering::Relaxed) + 1;
2797            if n % 100 == 0 {
2798                eprintln!(
2799                    "postprof: sync-wait {:.1} ms/ток | read_states {:.1} ms/ток ({n})",
2800                    SY.load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
2801                    RS.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
2802                );
2803            }
2804        }
2805        if let Some(rows) = lm_rows {
2806            crate::gpu::hostprof_encode_done(_mt0);
2807            let mut lg = attention::take_buf(rows.min(self.vocab_size));
2808            graph.read_logits(&mut lg);
2809            crate::gpu::hostprof_total(_mt0);
2810            lg.resize(self.vocab_size, 0.0);
2811            if let Some(c) = self.final_softcap {
2812                for l in lg.iter_mut() {
2813                    *l = c * (*l / c).tanh();
2814                }
2815            }
2816            self.graph_logits = Some(lg);
2817        }
2818        graph.read_h(h);
2819        if self.graph_head_required && self.graph_logits.is_none() {
2820            METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2821            self.fail_metal_graph("fused graph head completed without logits readback");
2822            return start;
2823        }
2824        METAL_GRAPH_TOK_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2825        METAL_GRAPH_LAYERS.fetch_add(
2826            end.saturating_sub(start) as u64,
2827            std::sync::atomic::Ordering::Relaxed,
2828        );
2829        if self.graph_head_required {
2830            METAL_GRAPH_HEAD_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2831        }
2832        // Device-attended layers: replay the CPU bookkeeping — append
2833        // the mirror's new K/V row (rope'd on the GPU) into the owner
2834        // cache, then bank this token's attention-importance mass.
2835        for li in dev_attn {
2836            let mut krow = attention::take_buf(nkv * hd);
2837            let mut vrow = attention::take_buf(nkv * hd);
2838            if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
2839                let cache = &mut self.kv_cache.layers[li];
2840                cache.append(&krow, &vrow, &[]);
2841                let n = cache.seq_len;
2842                let mut imp = attention::take_buf(n);
2843                crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
2844                cache.accumulate_imp(&imp);
2845                attention::recycle_buf(&mut imp);
2846            }
2847            attention::recycle_buf(&mut krow);
2848            attention::recycle_buf(&mut vrow);
2849        }
2850        if let Some((_, arm)) = ab {
2851            crate::gpu_metal::dense_ab_record(arm, _mt0.elapsed());
2852        }
2853        end
2854    }
2855
2856    pub fn new(
2857        tokenizer: Tokenizer,
2858        weights: PipelineWeights,
2859        hidden_size: usize,
2860        intermediate_size: usize,
2861        num_heads: usize,
2862        num_kv_heads: usize,
2863        head_dim: usize,
2864        num_layers: usize,
2865        physical_layers: usize,
2866        loop_final_norm: bool,
2867        vocab_size: usize,
2868        rms_eps: f64,
2869        rope_base: f32,
2870        norm_style: NormStyle,
2871        max_seq_len: usize,
2872        sampler_config: SamplerConfig,
2873    ) -> Self {
2874        let rng = match sampler_config.seed {
2875            Some(s) => SplitMix64::new(s),
2876            None => SplitMix64::from_entropy(),
2877        };
2878        let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
2879        let pool = Pool::from_env();
2880        if let Some(p) = &pool {
2881            tracing::info!("worker pool: {} threads", p.n_workers());
2882            // Keep the workers on the socket that holds the weights.
2883            if let Some(model) = weights
2884                .lm_head
2885                .model_arc()
2886                .or_else(|| weights.embed_tokens.model_arc())
2887            {
2888                let regions: Vec<&[u8]> =
2889                    model.tensors.iter().map(|t| model.entry_bytes(t)).collect();
2890                p.bind_numa(&regions);
2891            }
2892        }
2893        Self {
2894            gpu_plan: None,
2895            tokenizer: std::sync::Arc::new(tokenizer),
2896            kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
2897            sampler_config,
2898            weights,
2899            hidden_size,
2900            intermediate_size,
2901            num_heads,
2902            num_kv_heads,
2903            head_dim,
2904            num_layers,
2905            physical_layers,
2906            loop_final_norm,
2907            vocab_size,
2908            rms_eps,
2909            rope_base,
2910            norm_style,
2911            rotary_dim: head_dim,
2912            attention_heads_per_layer: None,
2913            kv_heads_per_layer: None,
2914            v_head_dim: None,
2915            layer_dump: std::env::var_os("CMF_LAYER_DUMP")
2916                .filter(|v| !v.is_empty())
2917                .map(std::path::PathBuf::from),
2918            graph_declines: std::cell::RefCell::new(Vec::new()),
2919            mimo_moe: Default::default(),
2920            vmf_cfg: None,
2921            gdn_cfg: None,
2922            kda_cfg: None,
2923            g3n: None,
2924            dsv4: None,
2925            dsv41: None,
2926            dsv41_vision: None,
2927            dsv41_prefill: None,
2928            qwen4_exp: None,
2929            dsv4_mtp: Vec::new(),
2930            dspark: None,
2931            dspark_pending: Vec::new(),
2932            dspark_hist: Vec::new(),
2933            dspark_real: Vec::new(),
2934            dspark_trunk_picks: Vec::new(),
2935            dspark_exp: Vec::new(),
2936            dspark_draft_ns: 0,
2937            logit_multiplier: None,
2938            cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
2939            graph_failed: std::sync::atomic::AtomicBool::new(false),
2940            kv_history: Vec::new(),
2941            kv_history_device: false,
2942            short_conv_cfg: None,
2943            mtp: None,
2944            mimo_mtp: None,
2945            verify_exact_moe: false,
2946            speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
2947            ignore_eos: false,
2948            draft_full_streak: 0,
2949            spec_k_adapt: None,
2950            spec_acc_ewma: 0.7,
2951            rng,
2952            sampler_scratch: SamplerScratch::default(),
2953            spec_forced: None,
2954            spec_q: Vec::new(),
2955            spec_p: Vec::new(),
2956            spec_res: Vec::new(),
2957            spec_qs: Vec::new(),
2958            spec_ps: Vec::new(),
2959            spec_ress: Vec::new(),
2960            mtp_graph_mode: None,
2961            #[cfg(target_os = "macos")]
2962            metal_verify: None,
2963            inv_freq,
2964            ws: ForwardScratch::new(hidden_size),
2965            pool,
2966            model: None,
2967            dyn_force_f32: false,
2968            dyn_skill_layers: Vec::new(),
2969            dyn_active: None,
2970            dyn_blend_loaded: false,
2971            dyn_phi_layer: None,
2972            dyn_phi_ema: Vec::new(),
2973            dyn_phi_seen: 0,
2974            dyn_router: None,
2975            o1_cfg: None,
2976            o1_epoch: 0,
2977            o1_flags: Vec::new(),
2978            trace: false,
2979            calib_temp: 1.0,
2980            confidence_on: true,
2981            embed_multiplier: 1.0,
2982            attn_scale: 1.0 / (head_dim as f32).sqrt(),
2983            swa: None,
2984            sliding_layers: None,
2985            anchor_core: None,
2986            bounded_rope: None,
2987            kv_prefix: KvPrefix::default(),
2988            last_prefill_tokens: 0,
2989            inv_freq_local: None,
2990            rotary_dim_local: None,
2991            rope_scale: 1.0,
2992            rope_scale_local: 1.0,
2993            global_attn: None,
2994            inv_freq_global: None,
2995            attn_v_norm: false,
2996            qk_norm_after_rope: false,
2997            final_softcap: None,
2998            head_clusters: None,
2999            attn_softcap: 0.0,
3000            graph_want_logits: false,
3001            graph_head_required: false,
3002            graph_logits: None,
3003            embryo_graph: None,
3004            graph_refused: std::sync::atomic::AtomicBool::new(false),
3005            graph_kv_id: {
3006                static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
3007                NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
3008            },
3009            #[cfg(test)]
3010            nll_test_fail_at: None,
3011            #[cfg(test)]
3012            nll_test_force_serial: false,
3013        }
3014    }
3015
3016    /// Enable/disable per-layer O(1) Nyström attention. Only Full
3017    /// layers are eligible (a linear layer keeps its own operator).
3018    /// Applies to generation (`generate*`/`forward_ids`): the prompt
3019    /// pass stays exact, then the state seals after prefill or at the
3020    /// deferred skeleton-safe boundary for short prompts; decode runs on
3021    /// the O(1) state. Teacher-forced scoring (`ppl_ids`) intentionally
3022    /// stays exact.
3023    pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
3024        if let Err(e) = self.try_set_o1(cfg) {
3025            tracing::error!("{e}");
3026        }
3027    }
3028
3029    /// True when the file carries a native bounded anchor
3030    /// (`arch.anchor_core`): its state is a fixed record the header
3031    /// fixes, and the post-hoc O(1) override is meaningless on it.
3032    pub fn bounded_native(&self) -> bool {
3033        self.anchor_core.is_some()
3034    }
3035
3036    /// Bytes the resident device graph holds for this pipeline's sequence:
3037    /// `(recurrent state, anchor KV/ring)`; None on the host path.
3038    pub fn device_state_bytes(&self) -> Option<(u64, u64)> {
3039        crate::gpu::embryo_device_state_bytes(self.graph_kv_id)
3040    }
3041
3042    /// Why an O(1) override is refused on this pipeline, if it is.
3043    pub fn o1_refusal(&self) -> Option<String> {
3044        self.anchor_core.as_ref().map(|ac| {
3045            format!(
3046                "--o1 / CMF_O1 refused: the anchor is native bounded \
3047                 (anchor_core kind={} window={} sink={}); the file's operator \
3048                 is executed as-is and no post-hoc Nyström overlay applies",
3049                ac.kind, ac.window, ac.sink
3050            )
3051        })
3052    }
3053
3054    /// `set_o1` that reports the refusal instead of logging it.
3055    pub fn try_set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) -> Result<(), String> {
3056        if let Some(c) = &cfg {
3057            if let Some(why) = self.o1_refusal() {
3058                self.o1_flags = Vec::new();
3059                self.o1_cfg = None;
3060                return Err(why);
3061            }
3062            if crate::nystrom::o1_deferred_boundary(c.w, c.sink).is_none() {
3063                self.o1_flags.clear();
3064                self.o1_cfg = None;
3065                return Err(format!(
3066                    "o1 disabled: w + sink + slack + 1 overflows usize (w={}, sink={})",
3067                    c.w, c.sink
3068                ));
3069            }
3070        }
3071        self.o1_flags = match &cfg {
3072            Some(c) => {
3073                let mut flags = c.layer_flags(self.num_layers);
3074                for (li, f) in flags.iter_mut().enumerate() {
3075                    // The Nyström state replaces a full-context plain
3076                    // softmax: a sliding window or a learned sink is not
3077                    // something it can represent, and a V narrower than
3078                    // the head is not what its streaming state stores.
3079                    // Those layers keep exact cache attention.
3080                    if *f
3081                        && (!matches!(
3082                            self.weights.layers[self.phys_layer(li)].attn,
3083                            AttnKind::Full { .. }
3084                        ) || self.layer_window(li).is_some()
3085                            || self.kv_cache.layers[li].sinks.is_some()
3086                            || self.layer_v_dim(li) != self.layer_geom(li).1)
3087                    {
3088                        *f = false;
3089                    }
3090                }
3091                flags
3092            }
3093            None => Vec::new(),
3094        };
3095        if let Some(c) = &cfg {
3096            let n = self.o1_flags.iter().filter(|&&f| f).count();
3097            tracing::info!(
3098                "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
3099                self.num_layers,
3100                c.m,
3101                c.w,
3102                c.sink,
3103                c.rect
3104            );
3105        }
3106        self.o1_cfg = cfg;
3107        Ok(())
3108    }
3109
3110    /// Install the file's bounded anchor: one fixed-size ring per
3111    /// `AttnKind::Bounded` layer (from the header, not per prompt) and
3112    /// the shared relative-rotation table. Must run after the RoPE setup
3113    /// (`set_rotary`, YaRN) so the table is built from the final
3114    /// `inv_freq`.
3115    pub fn install_bounded(
3116        &mut self,
3117        cfg: &cortiq_core::AnchorCoreConfig,
3118    ) -> Result<(), String> {
3119        if !cortiq_core::AnchorCoreConfig::KINDS.contains(&cfg.kind.as_str()) {
3120            return Err(format!(
3121                "anchor_core kind '{}' is not executable by this runtime",
3122                cfg.kind
3123            ));
3124        }
3125        if cfg.window == 0 {
3126            return Err("anchor_core.window must be >= 1".into());
3127        }
3128        let mut n = 0usize;
3129        for li in 0..self.num_layers {
3130            let pl = self.phys_layer(li);
3131            if let AttnKind::Bounded(w) = &self.weights.layers[pl].attn {
3132                if w.window != cfg.window || w.sink != cfg.sink {
3133                    return Err(format!(
3134                        "layer {li}: bounded weights (window {} sink {}) disagree with \
3135                         anchor_core (window {} sink {})",
3136                        w.window, w.sink, cfg.window, cfg.sink
3137                    ));
3138                }
3139                self.kv_cache.layers[li].install_bounded(cfg.window);
3140                n += 1;
3141            }
3142        }
3143        if n == 0 {
3144            return Err("anchor_core is present but no layer executes it".into());
3145        }
3146        let rope = crate::bounded::BoundedRope::new(cfg.window, &self.inv_freq, self.rope_scale);
3147        self.bounded_rope = Some(std::sync::Arc::new(rope));
3148        self.anchor_core = Some(cfg.clone());
3149        self.embryo_graph = None;
3150        tracing::info!(
3151            "bounded anchor {}: {n} layer(s), window {} sink {} — {} B of ring per layer",
3152            cfg.kind,
3153            cfg.window,
3154            cfg.sink,
3155            self.kv_cache.layers.iter().map(|l| l.bounded_state_bytes()).max().unwrap_or(0)
3156        );
3157        Ok(())
3158    }
3159
3160    /// Tag every layer's cache with its wire record kind and the model's
3161    /// operator identity (hash64 of `linear_core_identity` JSON) so the
3162    /// versioned state wire refuses a peer holding another operator.
3163    pub fn install_wire_identity(&mut self, identity: u64) {
3164        for li in 0..self.kv_cache.layers.len() {
3165            let pl = self.phys_layer(li);
3166            let kind = match self.weights.layers.get(pl).map(|l| &l.attn) {
3167                Some(AttnKind::Bounded(_)) => crate::kv_cache::WireKind::Bounded,
3168                Some(AttnKind::Linear(_))
3169                | Some(AttnKind::LinearGdn(_))
3170                | Some(AttnKind::ShortConv(_))
3171                | Some(AttnKind::Kda(_)) => crate::kv_cache::WireKind::Linear,
3172                _ => crate::kv_cache::WireKind::Full,
3173            };
3174            let l = &mut self.kv_cache.layers[li];
3175            l.wire_kind = kind;
3176            l.wire_identity = identity;
3177            // A per-layer geometry (`set_attn_geometry`, Gemma's global
3178            // heads) rebuilds a layer cache with index 0: re-tag it.
3179            l.wire_layer = li as u32;
3180        }
3181    }
3182
3183    /// Forget the reuse keys (legacy `kv_history` and the bounded prefix).
3184    pub fn clear_history(&mut self) {
3185        self.kv_history.clear();
3186        self.kv_history_device = false;
3187        self.kv_prefix.clear();
3188    }
3189
3190    /// Did THIS pipeline's token graph refuse for a structural reason?
3191    /// (Per pipeline: another lane's refusal, or a new pipeline of the
3192    /// same model, never changes it.)
3193    pub fn graph_refused(&self) -> bool {
3194        self.graph_refused
3195            .load(std::sync::atomic::Ordering::Relaxed)
3196    }
3197
3198    /// Remember a structural refusal of this pipeline's token graph.
3199    pub fn mark_graph_refused(&self) {
3200        if !self
3201            .graph_refused
3202            .swap(true, std::sync::atomic::Ordering::Relaxed)
3203        {
3204            tracing::info!(
3205                "token graph: unsupported for this pipeline (seq {}) — not retrying",
3206                self.graph_kv_id
3207            );
3208        }
3209    }
3210
3211    /// Position the resident Embryo graph holds for this pipeline's
3212    /// sequence (`Some(next position)`), `None` when the device holds no
3213    /// image of it (host-owned sequence, or none started).
3214    pub fn device_sequence_position(&self) -> Option<usize> {
3215        crate::gpu::embryo_device_next_position(self.graph_kv_id)
3216    }
3217
3218    /// Is the sequence this pipeline's reuse key describes owned by the
3219    /// device path it would take now? A prefix recorded on the resident
3220    /// graph continues only there (at exactly `n`), a host prefix only on
3221    /// the host; any mismatch means the next turn re-prefills from zero.
3222    fn prefix_owner_matches(&self, n: usize, recorded_on_device: bool) -> bool {
3223        let dev = self.device_sequence_position();
3224        if recorded_on_device {
3225            dev == Some(n) && self.embryo_resident_wanted()
3226        } else {
3227            dev.is_none()
3228        }
3229    }
3230
3231    /// The weights under the sequence changed (a real skill switch): every
3232    /// cached state was computed by other weights. Clear the host KV /
3233    /// ring / recurrent state AND the reuse keys — a surviving `kv_prefix`
3234    /// would let the next turn "extend" a prefix the new weights never
3235    /// saw — drop the packed resident graph (it holds the old FFN
3236    /// tensors; the next build packs the live ones under a fresh id) and
3237    /// reset its device sequence.
3238    pub(crate) fn invalidate_for_weight_change(&mut self) {
3239        self.clear_sequence_state();
3240        self.embryo_graph = None;
3241    }
3242
3243    /// Prompt positions the cache already holds when `input_ids`
3244    /// strictly EXTENDS the consumed prefix (0 otherwise). Bounded-native
3245    /// models answer from the fixed-size `kv_prefix` record; everything
3246    /// else from the legacy `kv_history` vector (which the network split
3247    /// also reads and writes).
3248    fn cached_prefix_len(&self, input_ids: &[u32]) -> usize {
3249        let (n, on_device) = if self.bounded_native() {
3250            (self.kv_prefix.extension(input_ids), self.kv_prefix.on_device())
3251        } else {
3252            let h = &self.kv_history;
3253            if !h.is_empty() && h.len() < input_ids.len() && input_ids[..h.len()] == h[..] {
3254                (h.len(), self.kv_history_device)
3255            } else {
3256                (0, false)
3257            }
3258        };
3259        // The owner tag: the state the key describes must live where this
3260        // turn will continue it (R4) — else a fresh sequence.
3261        if n > 0 && !self.prefix_owner_matches(n, on_device) {
3262            tracing::warn!(
3263                "kv-reuse refused: the cached prefix ({n} positions) was built on the {} path, \
3264                 the device now holds {:?} — re-prefilling from zero",
3265                if on_device { "resident device" } else { "host" },
3266                self.device_sequence_position()
3267            );
3268            return 0;
3269        }
3270        n
3271    }
3272
3273    /// Public view of the prefix-reuse decision for `input_ids` (positions
3274    /// the next `generate*` would take from the cache; 0 = fresh sequence).
3275    pub fn reusable_prefix_len(&self, input_ids: &[u32]) -> usize {
3276        self.cached_prefix_len(input_ids)
3277    }
3278
3279    /// Record the forwarded prefix as the next turn's reuse key. A
3280    /// bounded-native model extends the rolling record (`reused` = the
3281    /// positions this turn found cached); others keep the literal vector.
3282    /// Either way the key carries its OWNER: the resident device graph
3283    /// (it holds an image of this sequence) or the host.
3284    fn record_consumed_prefix(&mut self, consumed: &[u32], reused: usize) {
3285        let on_device = self.device_sequence_position().is_some();
3286        if self.bounded_native() {
3287            let keep = reused > 0 && reused == self.kv_prefix.len() && reused <= consumed.len();
3288            let prev_device = self.kv_prefix.on_device();
3289            self.kv_history.clear();
3290            self.kv_history_device = false;
3291            if keep && prev_device == on_device {
3292                self.kv_prefix.extend(&consumed[reused..]);
3293            } else {
3294                self.kv_prefix.set(consumed);
3295            }
3296            self.kv_prefix.set_on_device(on_device);
3297        } else {
3298            self.kv_history = consumed.to_vec();
3299            self.kv_history_device = on_device;
3300        }
3301    }
3302
3303    /// True when at least one layer runs the O(1) kernel.
3304    pub fn o1_active(&self) -> bool {
3305        self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
3306    }
3307
3308    /// Whether generation's prompt ingest is routed through the whole-token
3309    /// graph.  The bench uses this to label the measured generation prefill
3310    /// honestly; keep the predicate in Pipeline so CLI labels cannot drift
3311    /// from the production route.
3312    /// Positions per batched-graph submit for the prompt: `CMF_BATCH_K`
3313    /// when set (0 = one position at a time through the token graph),
3314    /// otherwise 32 on a discrete card whose prompt takes the graph route.
3315    /// The batched graph read a 2048-token prompt at 53 tok/s against 28.5
3316    /// one position at a time on an RTX PRO 4000 (Qwen3.8-27B q4tp: TTFT
3317    /// 39 s against 72), and its states are the speculative verify's,
3318    /// measured identical to the plain path. macOS keeps its own arm.
3319    pub fn generation_batch_k(&self) -> usize {
3320        if let Some(k) = std::env::var("CMF_BATCH_K")
3321            .ok()
3322            .and_then(|v| v.parse::<usize>().ok())
3323        {
3324            return k;
3325        }
3326        #[cfg(not(target_os = "macos"))]
3327        if self.graph_prefill_preferred() && !self.o1_active() {
3328            return 32;
3329        }
3330        0
3331    }
3332
3333    pub fn generation_graph_prefill(&self) -> bool {
3334        let graph = self.graph_prefill_preferred();
3335        // On wgpu, an active MTP head now consumes the trunk's graph batches
3336        // and warms its own block from those returned rows.  The selected
3337        // generation measurement is therefore the batched path, even though
3338        // the underlying GDN model still satisfies the graph-prefill
3339        // predicate.  Keep the CLI label tied to the actual route.  Native
3340        // Metal has a separate prefill-batch arm and retains its historical
3341        // label here.
3342        // A batched prompt (`generation_batch_k` > 0) is the batched graph
3343        // for every model on the graph route, not only those with an MTP
3344        // head — the label follows the route.
3345        #[cfg(not(target_os = "macos"))]
3346        if graph
3347            && self.generation_batch_k() > 0
3348            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
3349        {
3350            return false;
3351        }
3352        graph
3353    }
3354
3355    /// Device-side O(1) mirrors currently uploaded for this pipeline's
3356    /// sequence.  The count/bytes are zero before seal or after a fresh
3357    /// reset; callers use this to distinguish logical host state from the
3358    /// GPU allocation that actually serves decode.
3359    pub fn o1_device_stats(&self) -> (usize, u64) {
3360        crate::gpu::o1_device_stats(self.graph_kv_id)
3361    }
3362
3363    /// Arm query collection on the o1 layers (fresh prompt pass).
3364    /// Reset the o1 layers to Collecting for a fresh sequence. Pub for the
3365    /// network split: each side runs the o1 lifecycle over ITS OWN layers
3366    /// (begin before prefill, seal at the prefill barrier).
3367    pub fn o1_begin(&mut self) {
3368        self.o1_begin_with_prefix(None);
3369    }
3370
3371    /// Arm collection and optionally request a positive calibration prefix.
3372    /// The effective barrier is always at least the skeleton-safe floor, so
3373    /// a short requested prefix cannot create an exact-only runtime state.
3374    pub fn o1_begin_with_prefix(&mut self, requested_prefix: Option<usize>) {
3375        if let Some(c) = &self.o1_cfg {
3376            let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
3377            let boundary = requested_prefix.map(|p| {
3378                p.max(
3379                    crate::nystrom::o1_deferred_boundary(w, sink)
3380                        .expect("o1 config boundary validated in set_o1"),
3381                )
3382            });
3383            for (li, &f) in self.o1_flags.iter().enumerate() {
3384                if f {
3385                    self.kv_cache.layers[li].o1_begin_with_boundary(m, w, sink, rect, boundary);
3386                }
3387            }
3388        }
3389    }
3390
3391    /// Effective deferred boundary for a positive prefix request.
3392    fn o1_effective_boundary(&self, requested_prefix: usize) -> Option<usize> {
3393        self.o1_cfg.as_ref().and_then(|c| {
3394            crate::nystrom::o1_deferred_boundary(c.w, c.sink)
3395                .map(|floor| requested_prefix.max(floor))
3396        })
3397    }
3398
3399    fn o1_note_transition(&mut self) {
3400        // Drain every layer's one-shot bit before publishing one pipeline
3401        // epoch. `any()` would short-circuit on the first layer and leak the
3402        // remaining bits into later forwards, causing one epoch per layer.
3403        let mut transitioned = false;
3404        for (li, &flagged) in self.o1_flags.iter().enumerate() {
3405            if flagged {
3406                transitioned |= self.kv_cache.layers[li].take_o1_transition();
3407            }
3408        }
3409        if transitioned {
3410            self.o1_epoch = self.o1_epoch.wrapping_add(1);
3411        }
3412    }
3413
3414    fn o1_pending(&self) -> bool {
3415        self.o1_flags.iter().enumerate().any(|(li, &f)| {
3416            f && self.kv_cache.layers[li].seq_len > 0
3417                && self.kv_cache.layers[li].o1_pending_boundary().is_some()
3418        })
3419    }
3420
3421    fn o1_fail(&mut self, err: String) {
3422        tracing::error!("o1 deferred seal failed; terminating sequence: {err}");
3423        self.clear_sequence_state();
3424        self.graph_failed
3425            .store(true, std::sync::atomic::Ordering::Relaxed);
3426        self.cancel
3427            .store(true, std::sync::atomic::Ordering::Relaxed);
3428    }
3429
3430    /// Seal participating layers while retaining the exact state when the
3431    /// prompt is below the deferred boundary. A split worker may have
3432    /// collecting layers outside its owned span; zero-depth layers remain
3433    /// armed and are intentionally skipped until their peer runs them.
3434    pub fn o1_seal_checked(&mut self) -> Result<bool, String> {
3435        if self.o1_cfg.is_none() {
3436            return Ok(false);
3437        }
3438        let mut participating = false;
3439        for li in 0..self.num_layers {
3440            if !self.o1_flags.get(li).copied().unwrap_or(false) {
3441                continue;
3442            }
3443            if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3444                return Err(err);
3445            }
3446            if self.kv_cache.layers[li].seq_len == 0 {
3447                continue;
3448            }
3449            participating = true;
3450            let num_heads = self.layer_num_heads(li);
3451            self.kv_cache.layers[li].o1_seal_checked(num_heads)?;
3452        }
3453        self.o1_note_transition();
3454        for li in 0..self.num_layers {
3455            if self.o1_flags.get(li).copied().unwrap_or(false) {
3456                if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3457                    return Err(err);
3458                }
3459            }
3460        }
3461        Ok(participating
3462            && (0..self.num_layers).all(|li| {
3463                !self.o1_flags.get(li).copied().unwrap_or(false)
3464                    || self.kv_cache.layers[li].seq_len == 0
3465                    || self.kv_cache.layers[li].o1_sealed()
3466            }))
3467    }
3468
3469    /// Complete a deferred boundary after a full position/span forward.
3470    /// This is the pipeline owner for epoch publication and failure cleanup.
3471    fn o1_progress(&mut self) {
3472        if !self.o1_active() {
3473            return;
3474        }
3475        for li in 0..self.num_layers {
3476            if self.o1_flags.get(li).copied().unwrap_or(false) {
3477                if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3478                    self.o1_fail(err);
3479                    return;
3480                }
3481            }
3482        }
3483        // A qwen_attention row can seal in the middle of a complete layer
3484        // walk. Consume its transition even though the pending boundary has
3485        // already disappeared from the cache.
3486        self.o1_note_transition();
3487        if !self.o1_pending() {
3488            return;
3489        }
3490        if let Err(err) = self.o1_seal_checked() {
3491            self.o1_fail(err);
3492        }
3493    }
3494
3495    /// Turn a deferred O(1) failure raised by a hidden-only forward into the
3496    /// Result error its public batch/span caller must return. The failure
3497    /// path already cleared host/device sequence state; consume only the
3498    /// side-channel marker here and leave the pipeline reusable.
3499    fn check_o1_progress_failure(&mut self, phase: &str) -> Result<(), String> {
3500        if self
3501            .graph_failed
3502            .swap(false, std::sync::atomic::Ordering::Relaxed)
3503        {
3504            self.cancel
3505                .store(false, std::sync::atomic::Ordering::Relaxed);
3506            self.clear_sequence_state();
3507            return Err(format!("{phase}: deferred O(1) transition failed"));
3508        }
3509        Ok(())
3510    }
3511
3512    /// Freeze landmarks + skeleton state after the prompt pass and drop
3513    /// the o1 layers' full KV; decode then runs `step()` per token.
3514    /// Pub for the network split (see `o1_begin`).
3515    pub fn o1_seal(&mut self) {
3516        if let Err(err) = self.o1_seal_checked() {
3517            self.o1_fail(err);
3518        }
3519    }
3520
3521    /// Enable/disable the structured per-token telemetry trace (B4).
3522    pub fn set_trace(&mut self, on: bool) {
3523        self.trace = on;
3524    }
3525
3526    /// Replace all request-scoped sampler options and reset the random stream.
3527    /// This is required for deterministic `seed` semantics in pooled servers.
3528    pub fn set_sampler_config(&mut self, config: SamplerConfig) {
3529        self.rng = match config.seed {
3530            Some(seed) => SplitMix64::new(seed),
3531            None => SplitMix64::from_entropy(),
3532        };
3533        self.sampler_config = config;
3534    }
3535
3536    /// Toggle the per-token confidence reduction (a full-vocab
3537    /// softmax each token). `bench --core` turns it off so the timed
3538    /// loop matches llama-bench's core contract; the result's
3539    /// `confidence` vec is empty while off.
3540    pub fn set_confidence(&mut self, on: bool) {
3541        self.confidence_on = on;
3542    }
3543
3544    /// Set the confidence-calibration temperature (B1). Values ≤0 are
3545    /// clamped to raw (1.0).
3546    pub fn set_calib_temp(&mut self, t: f32) {
3547        self.calib_temp = if t > 1e-3 { t } else { 1.0 };
3548    }
3549
3550    /// The active calibration temperature (1.0 = raw probability).
3551    pub fn calib_temp(&self) -> f32 {
3552        self.calib_temp
3553    }
3554
3555    /// Partial rotary (Qwen3.5): rotate only the first `rotary_dim` dims;
3556    /// the frequency table is rebuilt over the rotary dims.
3557    pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
3558        self.rotary_dim = rotary_dim.min(self.head_dim);
3559        self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
3560        // The packed resident graph owns its own inverse-frequency plane;
3561        // changing RoPE after it was built must not leave a stale device
3562        // model behind the exact host configuration.
3563        self.embryo_graph = None;
3564    }
3565
3566    fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
3567        QwenAttnCfg {
3568            num_heads: self.num_heads,
3569            num_kv_heads: self.num_kv_heads,
3570            head_dim: self.head_dim,
3571            hidden_size: self.hidden_size,
3572            position,
3573            inv_freq: &self.inv_freq,
3574            rotary_dim: self.rotary_dim,
3575            scale: self.attn_scale,
3576            softcap: self.attn_softcap,
3577            window: None,
3578            v_norm: false,
3579            qk_norm_after_rope: self.qk_norm_after_rope,
3580            q_norm: None,
3581            k_norm: None,
3582            output_gate: false,
3583            softplus_gate: None,
3584            rope_scale: self.rope_scale,
3585            bias: None,
3586            rms_eps: self.rms_eps,
3587            norm_style: self.norm_style,
3588            pool: self.pool.as_deref(),
3589            v_head_dim: self.v_head_dim.unwrap_or(self.head_dim),
3590        }
3591    }
3592
3593    /// Generate text from a plain-text prompt. Streams tokens via `on_token`.
3594    pub fn generate(
3595        &mut self,
3596        prompt: &str,
3597        max_tokens: usize,
3598        task_mask: Option<&TaskMask>,
3599        on_token: Option<TokenCallback>,
3600    ) -> Result<GenerateResult, String> {
3601        let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
3602        self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
3603    }
3604
3605    /// Generate from a V4.1 multimodal prompt prepared by the vision module.
3606    /// Vision rows are encoded once and fed through the same bounded token walk as text.
3607    pub fn generate_from_vl(
3608        &mut self,
3609        input: &crate::dsv41_vision::PreparedVlInputs,
3610        max_tokens: usize,
3611        task_mask: Option<&TaskMask>,
3612        on_token: Option<TokenCallback>,
3613    ) -> Result<GenerateResult, String> {
3614        let Some(dsv41) = &self.dsv41 else {
3615            return Err("V4.1 multimodal input requires a DeepSeek-V4.1 pipeline".into());
3616        };
3617        if input.token_ids.is_empty() {
3618            return Err("empty V4.1 multimodal prompt".into());
3619        }
3620        if input.token_types.len() != input.token_ids.len() {
3621            return Err(format!(
3622                "V4.1 token type count {} != token count {}",
3623                input.token_types.len(),
3624                input.token_ids.len()
3625            ));
3626        }
3627        let dim = dsv41.2.dim;
3628        let mut embeddings = vec![None; input.token_ids.len()];
3629        let mut participates = vec![true; input.token_ids.len()];
3630        if !input.images.is_empty() {
3631            let vision = self
3632                .dsv41_vision
3633                .as_ref()
3634                .ok_or_else(|| "V4.1 image prompt has no loaded vision tower".to_string())?;
3635            for image in &input.images {
3636                let end = image.start.saturating_add(image.types.len());
3637                if end > input.token_ids.len() {
3638                    return Err(format!(
3639                        "V4.1 image span {}..{} exceeds prompt length {}",
3640                        image.start,
3641                        end,
3642                        input.token_ids.len()
3643                    ));
3644                }
3645                let mut span = vec![0.0f32; image.types.len() * dim];
3646                vision.fill_image_span(image, &mut span, self.pool.as_deref())?;
3647                for (offset, &kind) in image.types.iter().enumerate() {
3648                    let pos = image.start + offset;
3649                    if input.token_types[pos] != kind {
3650                        return Err(format!(
3651                            "V4.1 image type mismatch at position {pos}: {} != {kind}",
3652                            input.token_types[pos]
3653                        ));
3654                    }
3655                    embeddings[pos] = Some(span[offset * dim..(offset + 1) * dim].to_vec());
3656                    participates[pos] = false;
3657                }
3658            }
3659        }
3660        for (pos, &kind) in input.token_types.iter().enumerate() {
3661            if kind == crate::dsv41_vision::TEXT && embeddings[pos].is_some() {
3662                return Err(format!("V4.1 text position {pos} has an image embedding"));
3663            }
3664            if kind != crate::dsv41_vision::TEXT && embeddings[pos].is_none() {
3665                return Err(format!("V4.1 image position {pos} has no image embedding"));
3666            }
3667        }
3668        self.dsv41_prefill = Some((embeddings, participates));
3669        let result = self.generate_from_ids(&input.token_ids, max_tokens, task_mask, on_token);
3670        self.dsv41_prefill = None;
3671        result
3672    }
3673
3674    /// `None` when the mask forbids nothing (see `TaskMask::fully_open`).
3675    fn drop_open_mask<'m>(&self, m: Option<&'m TaskMask>) -> Option<&'m TaskMask> {
3676        m.filter(|m| !m.fully_open(self.intermediate_size, self.num_heads))
3677    }
3678
3679    /// Generate from prepared token ids (e.g. a chat template).
3680    ///
3681    /// With an MTP head, greedy generation without a task mask takes the
3682    /// speculative path: the MTP module drafts the token after next and
3683    /// the main model verifies both in one fused two-position forward
3684    /// (weights streamed once). The output is EXACTLY the vanilla greedy
3685    /// sequence — a rejected draft is rolled back — MTP only buys speed.
3686    pub fn generate_from_ids(
3687        &mut self,
3688        input_ids: &[u32],
3689        max_tokens: usize,
3690        task_mask: Option<&TaskMask>,
3691        on_token: Option<TokenCallback>,
3692    ) -> Result<GenerateResult, String> {
3693        self.generate_with_prompt_rows(input_ids, None, max_tokens, task_mask, on_token)
3694    }
3695
3696    /// Generate from complete prompt embeddings [token_count, hidden_size].
3697    /// Text rows can be obtained with `embed_id`; media rows replace only
3698    /// their expanded placeholder positions. Rows are already scaled and
3699    /// enter `PrefillIn::Hidden`, so a device graph must not re-embed them.
3700    /// Token-only KV reuse is disabled both into and out of this request.
3701    pub fn generate_from_embeds(
3702        &mut self,
3703        input_ids: &[u32],
3704        prompt_rows: &[f32],
3705        max_tokens: usize,
3706        task_mask: Option<&TaskMask>,
3707        on_token: Option<TokenCallback>,
3708    ) -> Result<GenerateResult, String> {
3709        if input_ids.is_empty()
3710            || input_ids.len().checked_mul(self.hidden_size) != Some(prompt_rows.len())
3711        {
3712            return Err("embedded prompt dimensions must be [tokens, hidden_size]".into());
3713        }
3714        if prompt_rows.iter().any(|x| !x.is_finite()) {
3715            return Err("embedded prompt contains non-finite values".into());
3716        }
3717        if !self.can_prefill_batched() || self.dyn_router.is_some()
3718            || self.o1_active() || self.mtp.is_some() || self.gpu_plan.is_some()
3719        {
3720            return Err("embedded prompts require the ordinary transformer path without O(1), dynamic routing, GPU splitting or a generic MTP head".into());
3721        }
3722        self.generate_with_prompt_rows(input_ids, Some(prompt_rows), max_tokens, task_mask, on_token)
3723    }
3724
3725    fn generate_with_prompt_rows(
3726        &mut self,
3727        input_ids: &[u32],
3728        prompt_rows: Option<&[f32]>,
3729        max_tokens: usize,
3730        task_mask: Option<&TaskMask>,
3731        mut on_token: Option<TokenCallback>,
3732    ) -> Result<GenerateResult, String> {
3733        #[cfg(target_os = "macos")]
3734        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
3735        if std::env::var("CMF_TRACE_H").is_ok() {
3736            eprintln!("input_ids: {input_ids:?}");
3737        }
3738        if input_ids.is_empty() {
3739            return Err("empty prompt: nothing to generate from".to_string());
3740        }
3741        // A prior graph failure is terminal for that sequence but must not
3742        // poison the next independent request.  Keep this flag separate from
3743        // the externally-owned cooperative cancel bit.
3744        self.graph_failed
3745            .store(false, std::sync::atomic::Ordering::Relaxed);
3746        // A mask that forbids nothing still costs every fused path and
3747        // whole-token graph, all of which are gated on `is_none()`. A
3748        // narrowed file whose one segment is always on carries exactly
3749        // such a mask — drop it here rather than pay 5x for a no-op.
3750        let task_mask = self.drop_open_mask(task_mask);
3751
3752        // Cross-turn KV reuse: a chat app resends the whole history
3753        // every turn; when the new ids strictly EXTEND what the cache
3754        // already holds, prefill only the tail — turn latency stays
3755        // proportional to the new text instead of the whole session.
3756        // Extension-only (no rollback), so it is exact for every layer
3757        // kind including recurrent state; MTP/o1/task-mask runs keep
3758        // the fresh-sequence path. CMF_KV_REUSE=0 disables.
3759        let mut reuse_from = {
3760            let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
3761            if on
3762                && prompt_rows.is_none()
3763                && task_mask.is_none()
3764                && self.mtp.is_none()
3765                && !(self.mimo_mtp.is_some() && self.speculative)
3766                && self.o1_cfg.is_none()
3767                && self.dsv41.is_none()
3768            {
3769                self.cached_prefix_len(input_ids)
3770            } else {
3771                0
3772            }
3773        };
3774        // The device may own rows the host tail prefill needs (wgpu decode
3775        // writes only its mirror): hand them to the host, or start fresh.
3776        if reuse_from > 0 && !self.prepare_kv_reuse(reuse_from) {
3777            reuse_from = 0;
3778        }
3779        self.last_prefill_tokens = input_ids.len() - reuse_from;
3780        let bounded_native = self.bounded_native();
3781        if reuse_from == 0 {
3782            // Fresh sequence — the cache holds absolute positions.
3783            self.clear_sequence_state();
3784        } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
3785            eprintln!(
3786                "kv-reuse: {} of {} prompt positions already cached",
3787                reuse_from,
3788                input_ids.len()
3789            );
3790        }
3791        crate::gpu::graph_race_begin_generation();
3792        // Optional bounded calibration prefix. Keep the requested value
3793        // even when it is longer than the prompt; the collecting layer will
3794        // defer at the effective boundary and remain exact for short input.
3795        let o1_prefill = if self.o1_active() && task_mask.is_none() {
3796            std::env::var("CMF_O1_PREFILL")
3797                .ok()
3798                .and_then(|v| v.parse::<usize>().ok())
3799                .filter(|&p| p > 0)
3800        } else {
3801            None
3802        };
3803        if task_mask.is_none() {
3804            self.o1_begin_with_prefix(o1_prefill);
3805        }
3806
3807        // Speculative decode is off under o1: a rejected draft can't be
3808        // rolled back out of the far accumulators / ring window (the
3809        // Nyström insertion is irreversible by design).
3810        // The wgpu token graph owns a device K/V mirror that speculative
3811        // rollback would desync — the two are mutually exclusive.
3812        let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
3813        // Graph speculative decode (`CMF_GRAPH_SPEC=1`): the MTP head
3814        // drafts, ONE batched graph submit verifies the whole chain.
3815        //
3816        // It now PAYS on Qwen3.6-27B / RTX 5090 — 51.1 tok/s against a
3817        // plain 49.4 at k=3, medians of three, 89% of drafts accepted,
3818        // and the greedy continuation is byte-identical to the plain
3819        // path. That took the batch matvec sharing its nibble unpack
3820        // across the batch (`CMF_MV_BK=2`); before it, the same round
3821        // measured 43.6, an 11% LOSS, which is what the earlier note
3822        // here described.
3823        //
3824        // Still opt-in. One model's win is not a default: the verify
3825        // rides `gdn_spec_restore` and a batched frame whose numerics
3826        // are the batch kernels', and that has to be shown on more than
3827        // one architecture before every greedy decode takes it.
3828        // Greedy (with or without penalties) verifies by argmax equality.
3829        // Sampling (temperature > 0) can go through speculative SAMPLING —
3830        // draft from the MTP head's own post-chain distribution, accept
3831        // with min(1, p/q), correct from max(0, p − q); the emitted stream
3832        // is distributed exactly as the plain sampler's — but it is
3833        // OPT-IN (`CMF_GRAPH_SPEC_SAMPLE=1`): measured on Qwen3.8-27B /
3834        // RTX 5090 at the instruct row (0.7 / 0.80 / 20 / presence 1.5)
3835        // it decoded 19-22 tok/s against a plain 40 — nine post-chain
3836        // distributions a round plus a lower acceptance than greedy's,
3837        // against a verify that costs 2.7 single tokens. The greedy arms
3838        // pay +10%; the sampling arm needs a cheaper verify first.
3839        // Native Metal HAS that verify: its eight-row tile is flat in b,
3840        // so a round costs ~1.9 plain tokens and the sampling arm pays at
3841        // 2.3 accepted per round — measured on Qwen3.8-27B q4tp / M4 at
3842        // the CLI defaults (0.7 / rep 1.1 / top-k 40, seed 42), a code
3843        // prompt: 9.0 tok/s against a plain 5.4 in the same window, and
3844        // the per-round watchdog turns it off where prose loses. So on
3845        // Metal the sampling arm is ON (`CMF_GRAPH_SPEC_SAMPLE=0` opts out)
3846        // — but only for a config the SPARSE chain serves (a top-k within
3847        // `sparse_ok`): without it a round builds nine 248k-float
3848        // distributions on the host, which is the 5090's measured loss and
3849        // not a cost the round-token proxy below can see. A top-k-less
3850        // sampling config keeps the plain path unless asked for by name.
3851        #[cfg(target_os = "macos")]
3852        let metal_graph = crate::gpu::q1_force()
3853            && crate::gpu::enabled_here()
3854            && std::env::var("CMF_GPU_BLOCK")
3855                .map(|v| v != "0")
3856                .unwrap_or(true);
3857        #[cfg(not(target_os = "macos"))]
3858        let metal_graph = false;
3859        let spec_sample_env = std::env::var("CMF_GRAPH_SPEC_SAMPLE").ok();
3860        // A round whose cost is the MEASURED one: greedy (argmax rows), or
3861        // sampling through the sparse chain. Anything else pays the dense
3862        // chain's host time, which no proxy can price.
3863        let spec_cheap_round = self.sampler_config.temperature < 1e-6
3864            || sampler::sparse_ok(&self.sampler_config);
3865        let spec_sampling_ok = self.sampler_config.temperature < 1e-6
3866            || match spec_sample_env.as_deref() {
3867                Some("1") => true,
3868                Some(_) => false,
3869                None => metal_graph && spec_cheap_round,
3870            };
3871        // ON by default for greedy on the wgpu graph: with the draft on
3872        // the graph and the verify bit-exact, it measured 58.7 tok/s
3873        // against a plain 48.1 on Qwen3.8-27B q4tp / RTX 5090 (k=4) and
3874        // 51.1 against 49.4 on Qwen3.6-27B, and a round that stops
3875        // paying turns itself off below (acceptance watchdog).
3876        // `CMF_GRAPH_SPEC=0` disables; `=1` was the old opt-in spelling.
3877        // …but only where the batched verify has its register-blocked
3878        // kernel: q4tp dense FFNs (graph kind 6). q4t and q8_2f verify
3879        // through tile GEMMs today and measured a LOSS (q8_2f 22 against
3880        // 29 tok/s), the 2-bit plane the same; those stay opt-in
3881        // (`CMF_GRAPH_SPEC=1`).
3882        // …at least in nine dense FFNs of ten: a healed file carries its
3883        // last two layers at q8_2f, and two tile-GEMM verifies among 64 do
3884        // not change the arithmetic (measured: the healed q4tp file
3885        // decodes at the plain file's rate and would otherwise sit out).
3886        let (mut dense_n, mut dense_q4tp) = (0usize, 0usize);
3887        for lw in &self.weights.layers {
3888            if let FfnKind::Dense(d) = &lw.ffn {
3889                dense_n += 1;
3890                if matches!(d.gate_proj.graph_weight(), Some((_, _, 6, _)))
3891                    && matches!(d.up_proj.graph_weight(), Some((_, _, 6, _)))
3892                    && matches!(d.down_proj.graph_weight(), Some((_, _, 6, _)))
3893                {
3894                    dense_q4tp += 1;
3895                }
3896            }
3897        }
3898        let spec_default_ok = dense_n == 0 || dense_q4tp * 10 >= dense_n * 9;
3899        // Penalties break the draft head's agreement with the trunk (a
3900        // 1.1 repetition penalty measured 2 of 16 accepted): not by
3901        // default there either — off Metal that rule is untouched, and
3902        // suppressed ids keep counting as a penalty there, because no
3903        // measurement on a discrete card says otherwise.
3904        //
3905        // On native Metal the penalized arms DO pay: the draft applies
3906        // the same penalty and the verify scores the penalized rows
3907        // exactly (`greedy_pen`, the plain loop's arithmetic), so the
3908        // text is the plain path's and only the round's shape changes.
3909        // Measured on this M4 — see the report for the interleaved run.
3910        let penalized = !metal_graph
3911            && (self.sampler_config.repetition_penalty != 1.0
3912                || self.sampler_config.presence_penalty != 0.0
3913                || !self.sampler_config.suppress_tokens.is_empty());
3914        // …and not on wgpu-over-Metal: the batched verify graph there
3915        // returned 0 accepted drafts and garbage text on a GDN hybrid
3916        // (16.08, Qwen3.5-0.8B) while Vulkan is bit-exact; the Mac's
3917        // default backend is native Metal without a batch graph anyway.
3918        #[cfg(feature = "gpu")]
3919        let metal_wgpu = graph_on && crate::gpu_wgpu::wgpu_backend_is_metal();
3920        #[cfg(not(feature = "gpu"))]
3921        let metal_wgpu = false;
3922        let spec_env = std::env::var("CMF_GRAPH_SPEC").ok();
3923        let spec_wanted = match spec_env.as_deref() {
3924            Some("0") => false,
3925            Some(_) => {
3926                if metal_wgpu {
3927                    tracing::warn!(
3928                        "CMF_GRAPH_SPEC forced on wgpu/Metal: the batched verify graph is not \
3929                         verified on this backend (garbage measured on Qwen3.5-0.8B)"
3930                    );
3931                }
3932                true
3933            }
3934            None => spec_default_ok && !penalized && !metal_wgpu,
3935        };
3936        // Native Metal: the b-row verify graph (`try_batch_graph_metal`)
3937        // stands where the wgpu batch graph stands on discrete cards
3938        // (`metal_graph`, above).
3939        let graph_spec = self.speculative
3940            && (graph_on || metal_graph)
3941            && self.mtp.is_some()
3942            && task_mask.is_none()
3943            && !self.o1_active()
3944            && spec_sampling_ok
3945            && spec_wanted;
3946        // Native Metal: say the route ONCE (RUST_LOG=info), so a user can
3947        // confirm the fast path without setting a single flag — every
3948        // knob below defaults to the measured-best value on the M4.
3949        #[cfg(target_os = "macos")]
3950        if metal_graph {
3951            static SAID: std::sync::Once = std::sync::Once::new();
3952            SAID.call_once(|| {
3953                let spec = if graph_spec {
3954                    let k = std::env::var("CMF_GRAPH_SPEC_K")
3955                        .ok()
3956                        .and_then(|v| v.parse::<usize>().ok())
3957                        .filter(|&v| (1..=8).contains(&v))
3958                        .unwrap_or(7);
3959                    let arm = if self.sampler_config.temperature < 1e-6 {
3960                        "greedy"
3961                    } else {
3962                        "sampling"
3963                    };
3964                    format!(
3965                        "spec k={k} {arm} (batched verify, draft shortlist {}, trial: proxy)",
3966                        Self::draft_vocab_rows(usize::MAX)
3967                    )
3968                } else if !self.speculative {
3969                    "spec off (CMF_MTP=0)".to_string()
3970                } else if self.mtp.is_none() {
3971                    "spec off (no MTP head)".to_string()
3972                } else if !spec_sampling_ok {
3973                    if spec_cheap_round {
3974                        "spec off (CMF_GRAPH_SPEC_SAMPLE=0)".to_string()
3975                    } else {
3976                        "spec off (sampling without a top-k: the dense chain \
3977                         costs more than it saves)"
3978                            .to_string()
3979                    }
3980                } else if !spec_wanted {
3981                    "spec off (CMF_GRAPH_SPEC=0 or non-q4tp FFNs)".to_string()
3982                } else if task_mask.is_some() {
3983                    "spec off (task mask)".to_string()
3984                } else {
3985                    "spec off (O(1) attention)".to_string()
3986                };
3987                let on = |var: &str| {
3988                    if std::env::var(var).as_deref() == Ok("0") {
3989                        "off"
3990                    } else {
3991                        "on"
3992                    }
3993                };
3994                tracing::info!(
3995                    "metal native: {spec}, state4 {}, async replay {}, prefill graph {}, \
3996                     MTP graph {}, attend {}, probe {}",
3997                    if crate::gpu_metal::state4_on() { "on" } else { "off" },
3998                    if crate::gpu_metal::async_replay_on() { "on" } else { "off" },
3999                    on("CMF_METAL_PREFILL"),
4000                    on("CMF_MTP_GRAPH"),
4001                    std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into()),
4002                    if crate::gpu::probe_enabled() { "bypassed (q1 force)" } else { "off" },
4003                );
4004            });
4005        }
4006        // GDN hybrids sit the fused-pair speculation out by default: the
4007        // recurrence is sequential, so the pair lane cannot parallelize
4008        // (the bench's own Pair line reads fused 1.28x TWO singles on the
4009        // 35B) and the draft's full-vocab head rides on top — measured 2x
4010        // SLOWER end to end (16.1 vs 32.4 tok/s on the 48-core stand).
4011        // CMF_MTP=1 forces it back for study.
4012        let pair_pays = self.gdn_cfg.is_none() || std::env::var("CMF_MTP").as_deref() == Ok("1");
4013        let spec_active = self.speculative
4014            && self.mtp.is_some()
4015            && task_mask.is_none()
4016            && !self.o1_active()
4017            && ((!graph_on && pair_pays && self.sampler_config.temperature < 1e-6) || graph_spec);
4018        // The MTP module is detached during generation so its mutable
4019        // state does not fight the borrow on `self`.
4020        let mut mtp = if spec_active { self.mtp.take() } else { None };
4021        if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
4022            eprintln!(
4023                "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
4024                mtp.is_some(),
4025                self.speculative,
4026                self.sampler_config.temperature < 1e-6,
4027            );
4028        }
4029        if let Some(m) = &mut mtp {
4030            m.kv.clear();
4031            // The MTP block's own device mirror starts over with its cache.
4032            crate::gpu::graph_kv_reset(self.mtp_kv_id());
4033            self.mtp_graph_mode = None;
4034        }
4035        // MiMo-V2's draft stack: greedy rounds (draft K with the chained
4036        // MTP layers, verify K+1 rows in one batched forward). Sampling
4037        // decodes plain; `CMF_MTP=0` / `CMF_MIMO_MTP=0` turn it off.
4038        let mimo_spec = self.speculative
4039            && self.mimo_mtp.is_some()
4040            && task_mask.is_none()
4041            && !self.o1_active()
4042            && self.dyn_router.is_none()
4043            && self.sampler_config.temperature < 1e-6
4044            && std::env::var("CMF_MIMO_MTP").as_deref() != Ok("0");
4045        if let Some(st) = self.mimo_mtp.as_mut() {
4046            st.reset();
4047            if mimo_spec && std::env::var_os("CMF_MIMO_MTP_PROBE").is_some() {
4048                Self::mimo_mtp_hist_cap(st, input_ids.len());
4049            }
4050        }
4051        // Dynamic router detached during decode (same borrow trick as MTP).
4052        // Speculative decode and dynamic routing are mutually exclusive
4053        // for now — the fused-pair path doesn't carry per-token φ.
4054        let mut router = if mtp.is_none() {
4055            self.dyn_router.take()
4056        } else {
4057            None
4058        };
4059        let mut reuse_from = reuse_from;
4060        if let Some(r) = &mut router {
4061            r.reset(); // active=backbone, matching a fresh overlay
4062            self.dyn_phi_seen = 0; // fresh φ EMA per generation
4063            if self.dyn_active.is_some() {
4064                // A real switch back to the backbone invalidates the
4065                // cache the reuse key was computed against.
4066                let _ = self.set_active_skill(None);
4067                reuse_from = 0;
4068                self.last_prefill_tokens = input_ids.len();
4069            }
4070        }
4071
4072        let mut all_ids = input_ids.to_vec();
4073        let mut generated = 0usize;
4074        let mut finish_reason = "max_tokens".to_string();
4075        let mut drafted = 0usize;
4076        let mut accepted = 0usize;
4077        // DeepSeek-V4's draft quality is strongly content-dependent.  Two
4078        // consecutive paid rounds with no extra token put it on a bounded
4079        // cooldown; predictable text keeps batching, ordinary prose falls
4080        // back to the exact walk instead of paying a slow draft forever.
4081        // Local to one generation so one difficult request cannot poison the
4082        // next one, and deliberately automatic — this is not a user knob.
4083        let mut dsv4_spec_bad = 0usize;
4084        let mut dsv4_spec_retry_at = 0usize;
4085        let mut confidence: Vec<f32> = Vec::new();
4086        let trace_on = self.trace;
4087        let calib_temp = self.calib_temp;
4088        let mut traces: Vec<TokenTrace> = Vec::new();
4089
4090        // ── Prefill: forward each prompt token once, KEEP the last hidden.
4091        //    Dense prefill runs in fused pairs (weights streamed once per
4092        //    two positions — bit-identical to sequential, proven by the
4093        //    pair tests). With MTP: warm the draft head on
4094        //    (hidden_p, token_{p+1}) pairs.
4095        let mut hidden = vec![0.0f32; self.hidden_size];
4096        let mut pos = reuse_from;
4097        // lm_head-in-graph is only sound when the very next logits
4098        // consumer is this loop's own (MTP and skill routing interleave
4099        // other forwards / can swap lm_head between forward and sample).
4100        // CMF_GPU_LMHEAD=0 keeps lm_head off the graph: the token reads back
4101        // the 8 KB hidden instead of ~1 MB of logits, and the head runs on
4102        // the host. A probe for how much of the graph's fixed per-token cost
4103        // is the logits readback (the layer sweep puts that fixed part at
4104        // 3.88 ms of an 18.5 ms frame).
4105        let fuse_lm = mtp.is_none()
4106            && router.is_none()
4107            && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
4108        self.graph_logits = None;
4109        self.graph_want_logits = false;
4110        let _tpf = std::time::Instant::now();
4111        let batch_k = self.generation_batch_k();
4112        if let Some(rows) = prompt_rows {
4113            let hs = self.hidden_size;
4114            let chunk = self.prefill_chunk().max(1);
4115            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4116                let end = (pos + chunk).min(input_ids.len());
4117                let hb = match self.prefill_input_rows(
4118                    PrefillIn::Hidden(&rows[pos * hs..end * hs]), pos, task_mask,
4119                ) {
4120                    Ok(hb) => hb,
4121                    Err(err) => {
4122                        self.finish_generation(&mut mtp, &mut router, true);
4123                        return Err(err);
4124                    }
4125                };
4126                if mimo_spec { self.mimo_note_rows(&hb, pos); }
4127                hidden.copy_from_slice(&hb[hb.len() - hs..]);
4128                pos = end;
4129            }
4130        }
4131        // DeepSeek-V4 owns a separate hyper-connection stack. Route it
4132        // before the generic prefill choices: those correctly reject an
4133        // empty `weights.layers`, but their final per-position fallback used
4134        // to consume the whole prompt before `dsv4::forward_chunk` could see
4135        // it. The batch implementation therefore existed without a live
4136        // production entry point.
4137        //
4138        // Bounded chunks preserve cancellation responsiveness. Only the
4139        // prompt's final chunk asks for logits; every earlier head projection
4140        // would produce 129 280 values that no caller reads.
4141        while self.qwen4_exp.is_some()
4142            && mtp.is_none()
4143            && pos < input_ids.len()
4144            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4145        {
4146            let token_id = input_ids[pos];
4147            let want_logits = pos + 1 == input_ids.len();
4148            let mut lg = Vec::new();
4149            if let Some(b) = &mut self.qwen4_exp {
4150                crate::qwen4_exp::forward_token(
4151                    &b.0,
4152                    &b.1,
4153                    &b.2,
4154                    &mut b.3,
4155                    token_id,
4156                    pos,
4157                    &self.inv_freq,
4158                    self.pool.as_deref(),
4159                    &mut lg,
4160                    want_logits,
4161                );
4162            }
4163            if want_logits {
4164                self.graph_logits = Some(lg);
4165            }
4166            pos += 1;
4167            hidden.fill(0.0);
4168        }
4169        while self.dsv4.is_some()
4170            && mtp.is_none()
4171            && pos < input_ids.len()
4172            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4173        {
4174            let end = (pos + prefill_chunk()).min(input_ids.len());
4175            let ids: Vec<u32> = input_ids[pos..end].to_vec();
4176            let mut lg = Vec::new();
4177            if let Some(b) = &mut self.dsv4 {
4178                let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
4179                crate::dsv4::forward_chunk(
4180                    g,
4181                    layers,
4182                    &cfg,
4183                    st,
4184                    &ids,
4185                    pos,
4186                    &self.inv_freq,
4187                    self.pool.as_deref(),
4188                    &mut lg,
4189                    end == input_ids.len(),
4190                );
4191            }
4192            if end == input_ids.len() {
4193                self.graph_logits = Some(lg);
4194            }
4195            pos = end;
4196            hidden = vec![0.0; self.hidden_size];
4197        }
4198        let dsv41_prefill = self.dsv41_prefill.take();
4199        while self.dsv41.is_some()
4200            && mtp.is_none()
4201            && pos < input_ids.len()
4202            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4203        {
4204            let end = (pos + prefill_chunk()).min(input_ids.len());
4205            let ids: Vec<u32> = input_ids[pos..end].to_vec();
4206            let mut lg = Vec::new();
4207            if let Some(b) = &mut self.dsv41 {
4208                let (g, layers, cfg, st) = (&b.0, &b.1, &b.2, &mut b.3);
4209                if let Some((embeddings, participates)) = dsv41_prefill.as_ref() {
4210                    crate::dsv41::forward_chunk_masked_with_embeddings(
4211                        g,
4212                        layers,
4213                        cfg,
4214                        st,
4215                        &ids,
4216                        pos,
4217                        &embeddings[pos..end],
4218                        &participates[pos..end],
4219                        self.pool.as_deref(),
4220                        &mut lg,
4221                    );
4222                } else {
4223                    crate::dsv41::forward_chunk(
4224                        g,
4225                        layers,
4226                        cfg,
4227                        st,
4228                        &ids,
4229                        pos,
4230                        self.pool.as_deref(),
4231                        &mut lg,
4232                    );
4233                }
4234            }
4235            if end == input_ids.len() {
4236                self.graph_logits = Some(lg);
4237            }
4238            pos = end;
4239            hidden = vec![0.0; self.hidden_size];
4240        }
4241        // With dynamic routing, prefill sequentially so the φ hook fires
4242        // over the PROMPT — the router enters decode with a warm φ (the
4243        // fused-pair path skips the per-layer φ capture). o1 layers
4244        // collect their query trace in both the single and pair paths.
4245        let dyn_prefill = router.is_some();
4246        // Optional bounded calibration prefix for generation.  The normal
4247        // O(1) path seals after the full prompt; this explicit knob instead
4248        // runs only the requested prefix through exact attention, seals the
4249        // Nyström state, and streams the rest of the prompt through the same
4250        // O(1) step used by decode.  It keeps the O(1) layers' Q trace and
4251        // temporary full KV bounded by the prefix while leaving the default
4252        // full-prompt quality profile untouched.
4253        let o1_prefill_limit = o1_prefill
4254            .and_then(|requested| self.o1_effective_boundary(requested))
4255            .map(|boundary| boundary.min(input_ids.len()));
4256        let mut o1_sealed = false;
4257        if let Some(limit) = o1_prefill_limit {
4258            // Reuse the exact batched prefix machinery when available; it
4259            // records the same per-position Q trace as the full prefill.
4260            if self.can_prefill_batched() && limit > 2 {
4261                let chunk = self.prefill_chunk();
4262                let hs = self.hidden_size;
4263                while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4264                    let end = (pos + chunk).min(limit);
4265                    let hb = self.prefill_batch(&input_ids[pos..end], pos);
4266                    hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4267                    pos = end;
4268                }
4269            } else {
4270                while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4271                    hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, None);
4272                    pos += 1;
4273                }
4274            }
4275            if pos >= limit {
4276                o1_sealed = match self.o1_seal_checked() {
4277                    Ok(sealed) => sealed,
4278                    Err(err) => {
4279                        self.finish_generation(&mut mtp, &mut router, true);
4280                        return Err(err);
4281                    }
4282                };
4283                tracing::info!(
4284                    "o1 bounded prompt prefix: requested={} effective={} processed={} of {} token(s)",
4285                    o1_prefill.unwrap_or(0),
4286                    self.o1_effective_boundary(o1_prefill.unwrap_or(0))
4287                        .unwrap_or(limit),
4288                    limit,
4289                    input_ids.len()
4290                );
4291            }
4292        }
4293        // q1 hybrids on Metal: the per-position GPU token graph beats
4294        // the CPU chunk-GEMM (whose wall is the sequential scalar GDN
4295        // recurrence), so prefill goes position-by-position through the
4296        // same graph as decode. Pure-attention models keep the batched
4297        // path — there the chunk-GEMM amortization wins.
4298        let graph_prefill = self.graph_prefill_preferred();
4299        // Native Metal, q4tp GDN hybrids: the prompt through the b-row
4300        // rows graph — projections as GEMMs over up to 512 positions, the
4301        // GDN recurrence in registers on the device, K/V rows appended by
4302        // the chunk — instead of one token-graph submit per position (the
4303        // 27B: 8 tok/s → GEMM-bound). The MTP warm-up rows come out of one
4304        // batched run of the block per chunk. Any refusal leaves the rest
4305        // of the prompt to the sequential paths below.
4306        #[cfg(target_os = "macos")]
4307        if task_mask.is_none()
4308            && !dyn_prefill
4309            && (crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in())
4310            && crate::gpu::enabled_here()
4311            && self.gdn_cfg.is_some()
4312            && self.g3n.is_none()
4313            && input_ids.len() > 8
4314            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
4315            && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
4316        {
4317            let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
4318                .ok()
4319                .and_then(|v| v.parse().ok())
4320                .filter(|&v| (16..=512).contains(&v))
4321                .unwrap_or(256);
4322            let hs = self.hidden_size;
4323            let _tp = std::time::Instant::now();
4324            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4325                let end = (pos + chunk).min(input_ids.len());
4326                let hb = match self.prefill_batch_metal(&input_ids[pos..end], pos) {
4327                    MetalPrefillOutcome::Completed(hb) => hb,
4328                    MetalPrefillOutcome::Declined => break,
4329                    MetalPrefillOutcome::Failed => {
4330                        self.finish_generation(&mut mtp, &mut router, true);
4331                        return Err("ordinary Metal prefill failed after admission".into());
4332                    }
4333                };
4334                if let Some(m) = &mut mtp {
4335                    let n_pairs = if end < input_ids.len() {
4336                        end - pos
4337                    } else {
4338                        end - pos - 1
4339                    };
4340                    if n_pairs > 0 {
4341                        let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
4342                            .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
4343                            .collect();
4344                        if !self.mtp_warm_batch_metal(m, &pairs, pos) {
4345                            for (j, (h, t)) in pairs.iter().enumerate() {
4346                                let h = h.to_vec();
4347                                let _ = self.mtp_step(m, &h, *t, pos + j);
4348                            }
4349                        }
4350                    }
4351                }
4352                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4353                pos = end;
4354            }
4355            if std::env::var("CMF_PREFILL_PROF").is_ok() {
4356                eprintln!(
4357                    "metal-prefill: {} of {} tokens in {:.1} ms",
4358                    pos,
4359                    input_ids.len(),
4360                    _tp.elapsed().as_secs_f64() * 1e3
4361                );
4362            }
4363        }
4364        self.mimo_moe_prepare();
4365        // A MoE stack larger than the card (MiMo-V2 q4tp on 96 GB): the
4366        // batched wgpu graph runs the device prefix of every chunk — its
4367        // experts resident — and the host's batched layer walk finishes
4368        // the chunk. Any refusal leaves the rest of the prompt to the
4369        // chunked prefill below.
4370        #[cfg(not(target_os = "macos"))]
4371        if task_mask.is_none()
4372            && !dyn_prefill
4373            && !graph_prefill
4374            && mtp.is_none()
4375            && o1_prefill.is_none()
4376            && !self.o1_active()
4377            && input_ids.len() > 2
4378            && self.batch_prefix_prefill()
4379        {
4380            let chunk = self.prefill_chunk().max(1);
4381            let hs = self.hidden_size;
4382            let t_bp = std::time::Instant::now();
4383            let pos0 = pos;
4384            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4385                let end = (pos + chunk).min(input_ids.len());
4386                let bk = end - pos;
4387                let mut hiddens = vec![0f32; bk * hs];
4388                for (j, &id) in input_ids[pos..end].iter().enumerate() {
4389                    hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4390                }
4391                let positions: Vec<usize> = (pos..end).collect();
4392                let mut run = 0usize;
4393                let outcome = self.try_batch_graph_wgpu_prefix(
4394                    &mut hiddens,
4395                    &positions,
4396                    bk,
4397                    None,
4398                    Some(&mut run),
4399                );
4400                match outcome {
4401                    crate::gpu::BatchGraphOutcome::Completed => {
4402                        let hb = if run < self.num_layers {
4403                            self.prefill_batch_span(
4404                                PrefillIn::Hidden(&hiddens),
4405                                pos,
4406                                None,
4407                                run,
4408                                self.num_layers,
4409                            )
4410                        } else {
4411                            hiddens
4412                        };
4413                        if mimo_spec {
4414                            self.mimo_note_rows(&hb, pos);
4415                        }
4416                        hidden.copy_from_slice(&hb[(bk - 1) * hs..]);
4417                        pos = end;
4418                    }
4419                    crate::gpu::BatchGraphOutcome::Failed => {
4420                        self.finish_generation(&mut mtp, &mut router, true);
4421                        return Err("batched prefix prefill failed after admission".into());
4422                    }
4423                    crate::gpu::BatchGraphOutcome::Declined => {
4424                        // Earlier chunks left their prefix rows on the
4425                        // device only: the host walk below needs them.
4426                        #[cfg(feature = "gpu")]
4427                        if pos > pos0 {
4428                            self.pull_lagging_host_kv(0, self.num_layers, pos);
4429                        }
4430                        break;
4431                    }
4432                }
4433            }
4434            if std::env::var("CMF_PREFILL_PROF").is_ok() {
4435                eprintln!(
4436                    "batch-prefix prefill: {} of {} tokens in {:.1} ms",
4437                    pos - pos0,
4438                    input_ids.len(),
4439                    t_bp.elapsed().as_secs_f64() * 1e3
4440                );
4441            }
4442        }
4443        if task_mask.is_none()
4444            && !dyn_prefill
4445            && !graph_prefill
4446            && self.can_prefill_batched()
4447            && self.g3n.is_none()
4448            && o1_prefill.is_none()
4449            && input_ids.len() > 2
4450        {
4451            // Production prefill = the same chunked prefill-GEMM that
4452            // bench/PPL measure (roadmap §3 P0: generation used to warm
4453            // the prompt with the slower pair path — the published
4454            // prefill number didn't match real TTFT). MTP warm-up reads
4455            // each position's hidden straight from the chunk result.
4456            let chunk = self.prefill_chunk();
4457            let hs = self.hidden_size;
4458            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4459                let end = (pos + chunk).min(input_ids.len());
4460                let hb = self.prefill_batch(&input_ids[pos..end], pos);
4461                if mimo_spec {
4462                    self.mimo_note_rows(&hb, pos);
4463                }
4464                if let Some(m) = &mut mtp {
4465                    let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4466                        .ok()
4467                        .and_then(|v| v.parse().ok())
4468                        .unwrap_or(0);
4469                    for p in pos..end {
4470                        if p + 1 < input_ids.len() {
4471                            if probe >= 1 && p + 2 < input_ids.len() {
4472                                // Teacher-forced chain acceptance (see the
4473                                // tail loop's twin): the warm-up row stays,
4474                                // the chain's rows roll back.
4475                                let (d1, mut hx) = self.mtp_step_h(
4476                                    m,
4477                                    &hb[(p - pos) * hs..(p - pos + 1) * hs],
4478                                    input_ids[p + 1],
4479                                    p,
4480                                );
4481                                let mut ok = d1 == input_ids[p + 2];
4482                                Self::chain_probe_note(0, ok);
4483                                let mut d_prev = d1;
4484                                let mut extra = 0usize;
4485                                for j in 1..probe {
4486                                    if p + 2 + j >= input_ids.len() {
4487                                        break;
4488                                    }
4489                                    let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
4490                                    extra += 1;
4491                                    ok = ok && dj == input_ids[p + 2 + j];
4492                                    Self::chain_probe_note(j, ok);
4493                                    d_prev = dj;
4494                                    hx = hj;
4495                                }
4496                                m.kv.truncate_last(extra);
4497                            } else {
4498                                let _ = self.mtp_step(
4499                                    m,
4500                                    &hb[(p - pos) * hs..(p - pos + 1) * hs],
4501                                    input_ids[p + 1],
4502                                    p,
4503                                );
4504                            }
4505                        }
4506                    }
4507                }
4508                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4509                pos = end;
4510            }
4511        }
4512        let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
4513        if task_mask.is_none()
4514            && !dyn_prefill
4515            && !graph_prefill
4516            && !pair_off
4517            && self.pair_supported()
4518            && o1_prefill.is_none()
4519        {
4520            while pos + 1 < input_ids.len()
4521                && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4522            {
4523                let e1 = self.embed_single(input_ids[pos]);
4524                let e2 = self.embed_single(input_ids[pos + 1]);
4525                let (h1, h2) = self.forward_pair(&e1, &e2, pos);
4526                if mimo_spec {
4527                    self.mimo_note_rows(&h1, pos);
4528                    self.mimo_note_rows(&h2, pos + 1);
4529                }
4530                // Both prefill tokens are real → commit lane-2 states.
4531                self.commit_linear_scratch();
4532                if let Some(m) = &mut mtp {
4533                    let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
4534                    if pos + 2 < input_ids.len() {
4535                        let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4536                            .ok()
4537                            .and_then(|v| v.parse().ok())
4538                            .unwrap_or(0);
4539                        if probe >= 1 && pos + 3 < input_ids.len() {
4540                            // Same teacher-forced chain table as the tail
4541                            // loop below, fed from the pair path that owns
4542                            // most prefill positions.
4543                            let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
4544                            let mut ok = d1 == input_ids[pos + 3];
4545                            Self::chain_probe_note(0, ok);
4546                            let mut d_prev = d1;
4547                            let mut extra = 0usize;
4548                            for j in 1..probe {
4549                                if pos + 3 + j >= input_ids.len() {
4550                                    break;
4551                                }
4552                                let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
4553                                extra += 1;
4554                                ok = ok && dj == input_ids[pos + 3 + j];
4555                                Self::chain_probe_note(j, ok);
4556                                d_prev = dj;
4557                                hx = hj;
4558                            }
4559                            m.kv.truncate_last(extra);
4560                        } else {
4561                            let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
4562                        }
4563                    }
4564                }
4565                hidden = h2;
4566                pos += 2;
4567            }
4568        }
4569        // Batched GPU prefill for the wgpu decode graph (GDN hybrids): K prompt
4570        // positions per submit — projections/FFN as GEMMs (weight once per K),
4571        // attention/GDN looped inside — instead of one whole-graph submit per
4572        // position. Falls through to the per-position graph on any refusal.
4573        // Batched prefill is opt-in (CMF_BATCH_K>0). Default 0 = per-position
4574        // graph prefill. (Steady-state decode is provably identical either way —
4575        // token-graph submit and lm_head both unchanged — so this only trades
4576        // prefill wall.)
4577        // A bounded O(1) prefix is the one post-seal prompt interval: only
4578        // admit its batch when the device O(1) route is explicitly enabled and
4579        // every sealed layer exposes a portable view. The same batch size and
4580        // refusal behavior remain the ordinary controls/comparator.
4581        let o1_batch_ready = o1_sealed
4582            && o1_prefill.is_some()
4583            && mtp.is_none()
4584            && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
4585            && (0..self.num_layers).all(|li| {
4586                let cache = &self.kv_cache.layers[self.phys_layer(li)];
4587                cache.o1.is_none() || cache.o1_views().is_some()
4588            });
4589        // The ordinary graph-prefill route can share each completed trunk
4590        // chunk with an attached MTP head.  Keep chain probing on its
4591        // established per-position path: the probe deliberately needs every
4592        // teacher-forced draft row and its rollback table.
4593        let mtp_batch_prefill = mtp.is_some()
4594            && graph_prefill
4595            && task_mask.is_none()
4596            && !dyn_prefill
4597            && !self.o1_active()
4598            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err();
4599        if batch_k > 0
4600            && (graph_prefill || o1_batch_ready)
4601            && task_mask.is_none()
4602            && (!self.o1_active() || o1_batch_ready)
4603            && (mtp.is_none() || mtp_batch_prefill)
4604            && !dyn_prefill
4605            && pos + 1 < input_ids.len()
4606        {
4607            let hs = self.hidden_size;
4608            let chunk = batch_k;
4609            while pos < input_ids.len() {
4610                let end = (pos + chunk).min(input_ids.len());
4611                let bk = end - pos;
4612                let mut hiddens = vec![0f32; bk * hs];
4613                for (j, &id) in input_ids[pos..end].iter().enumerate() {
4614                    hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4615                }
4616                let positions: Vec<usize> = (pos..end).collect();
4617                let t_chunk = std::time::Instant::now();
4618                let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
4619                let ok_b = outcome == crate::gpu::BatchGraphOutcome::Completed;
4620                if std::env::var("CMF_GRAPH_PROF").is_ok() {
4621                    let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
4622                    eprintln!(
4623                        "batch-chunk: phase=prompt mode={} k={bk} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
4624                        if o1_batch_ready {
4625                            "o1"
4626                        } else if mtp_batch_prefill {
4627                            "ordinary_mtp"
4628                        } else {
4629                            "ordinary"
4630                        },
4631                        bk as f64 / (ms / 1000.0)
4632                    );
4633                }
4634                {
4635                    use std::sync::atomic::{AtomicBool, Ordering};
4636                    static SAID: AtomicBool = AtomicBool::new(false);
4637                    if !SAID.swap(true, Ordering::Relaxed) {
4638                        if ok_b {
4639                            tracing::info!(
4640                                "batched prefill: ACTIVE mode={} (k={bk})",
4641                                if o1_batch_ready {
4642                                    "o1"
4643                                } else if mtp_batch_prefill {
4644                                    "ordinary_mtp"
4645                                } else {
4646                                    "ordinary"
4647                                }
4648                            );
4649                        } else {
4650                            tracing::warn!("batched prefill {:?} — per-position graph", outcome);
4651                        }
4652                    }
4653                }
4654                if ok_b {
4655                    if mimo_spec {
4656                        self.mimo_note_rows(&hiddens, pos);
4657                    }
4658                    if mtp_batch_prefill {
4659                        let n_pairs = mtp_prefill_pair_count(pos, end, input_ids.len());
4660                        if n_pairs > 0 {
4661                            // `hiddens` is owned by this chunk, so materialize
4662                            // row slices before borrowing the detached MTP
4663                            // module.  The last prompt row has no successor;
4664                            // the helper above is the single source of that
4665                            // boundary rule.
4666                            let rows: Vec<Vec<f32>> = (0..n_pairs)
4667                                .map(|j| hiddens[j * hs..(j + 1) * hs].to_vec())
4668                                .collect();
4669                            let pairs: Vec<(&[f32], u32)> = rows
4670                                .iter()
4671                                .enumerate()
4672                                .map(|(j, row)| (row.as_slice(), input_ids[pos + j + 1]))
4673                                .collect();
4674                            if std::env::var("CMF_GRAPH_PROF").is_ok() {
4675                                eprintln!(
4676                                    "mtp-warm: phase=prompt mode=ordinary_mtp first_pos={} pairs={} last_pos={}",
4677                                    pos,
4678                                    n_pairs,
4679                                    pos + n_pairs - 1,
4680                                );
4681                            }
4682                            let warm_error = if let Some(m) = mtp.as_mut() {
4683                                self.mtp_warm_prefill_pairs(m, &pairs, pos).err()
4684                            } else {
4685                                None
4686                            };
4687                            if let Some(err) = warm_error {
4688                                // The trunk batch was already admitted.  A
4689                                // failed MTP warm-up therefore clears both
4690                                // mirrors and exits; continuing would pair a
4691                                // current trunk state with a stale MTP cache.
4692                                self.finish_generation(&mut mtp, &mut router, true);
4693                                return Err(err.to_string());
4694                            }
4695                        }
4696                    }
4697                    hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
4698                    pos = end;
4699                } else if outcome == crate::gpu::BatchGraphOutcome::Failed {
4700                    // A failed batch may have advanced a device recurrent
4701                    // state (ordinary GDN or sealed O(1)). A CPU fallback
4702                    // would then observe stale accumulators, so clear the
4703                    // request state and make the failure explicit.
4704                    self.finish_generation(&mut mtp, &mut router, true);
4705                    return Err(if o1_batch_ready {
4706                        "sealed O(1) batch graph failed after admission".to_string()
4707                    } else {
4708                        "ordinary recurrent batch graph failed after admission".to_string()
4709                    });
4710                } else {
4711                    break; // unsupported → per-position graph handles the rest
4712                }
4713            }
4714        }
4715        // Resident Embryo graph: the prompt in chunks of one submit each
4716        // instead of one whole-graph submit per position; the last chunk
4717        // carries the logits exactly as the per-position walk would.
4718        if graph_prefill
4719            && task_mask.is_none()
4720            && mtp.is_none()
4721            && !dyn_prefill
4722            && pos == 0
4723            && input_ids.len() > 1
4724            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4725        {
4726            if let Some(lg) = self.embryo_prefill_chunked(input_ids, 0) {
4727                self.graph_logits = Some(lg);
4728                hidden = vec![0.0; self.hidden_size];
4729                pos = input_ids.len();
4730            }
4731        }
4732        while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4733            self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
4734            hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
4735            if mimo_spec {
4736                self.mimo_note_rows(&hidden, pos);
4737            }
4738            if let Some(m) = &mut mtp {
4739                if pos + 1 < input_ids.len() {
4740                    // `CMF_MTP_CHAIN_PROBE=k`: teacher-forced acceptance of a
4741                    // CHAINED draft — iterate the head on its own hidden k
4742                    // deep and score every depth against the prompt's real
4743                    // continuation. The economics of a k-token speculative
4744                    // round stand or fall on this table.
4745                    let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4746                        .ok()
4747                        .and_then(|v| v.parse().ok())
4748                        .unwrap_or(0);
4749                    if probe >= 1 && pos + 2 < input_ids.len() {
4750                        let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
4751                        let mut ok = d1 == input_ids[pos + 2];
4752                        Self::chain_probe_note(0, ok);
4753                        let mut d_prev = d1;
4754                        let mut extra = 0usize;
4755                        for j in 1..probe {
4756                            if pos + 2 + j >= input_ids.len() {
4757                                break;
4758                            }
4759                            let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
4760                            extra += 1;
4761                            ok = ok && dj == input_ids[pos + 2 + j];
4762                            Self::chain_probe_note(j, ok);
4763                            d_prev = dj;
4764                            hx = hj;
4765                        }
4766                        // The chain's rows are speculation, not the prompt —
4767                        // keep only the warmup row the plain path would add.
4768                        m.kv.truncate_last(extra);
4769                    } else {
4770                        let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
4771                    }
4772                }
4773            }
4774            pos += 1;
4775        }
4776        if std::env::var("CMF_PREFILL_PROF").is_ok() {
4777            eprintln!(
4778                "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
4779                input_ids.len(),
4780                _tpf.elapsed().as_secs_f64() * 1000.0
4781            );
4782        }
4783        if self
4784            .graph_failed
4785            .swap(false, std::sync::atomic::Ordering::Relaxed)
4786        {
4787            // MTP is detached for speculative generation.  Restore the
4788            // module before returning the terminal graph error; otherwise a
4789            // failed request would silently remove the head from a pooled
4790            // pipeline and the next request would lose its configured route.
4791            self.finish_generation(&mut mtp, &mut router, true);
4792            return Err("GPU token graph failed during prefill".to_string());
4793        }
4794        // Cancelled mid-prefill: the cache holds a partial prompt —
4795        // drop the reuse history and return an empty generation.
4796        if self
4797            .cancel
4798            .swap(false, std::sync::atomic::Ordering::Relaxed)
4799        {
4800            // A cancelled prefill can already have advanced the device
4801            // mirror. Drop the whole partial sequence so a pooled pipeline
4802            // cannot carry that state into its next request.
4803            self.finish_generation(&mut mtp, &mut router, true);
4804            return Ok(GenerateResult {
4805                text: String::new(),
4806                token_ids: Vec::new(),
4807                prompt_tokens: input_ids.len(),
4808                tokens_generated: 0,
4809                finish_reason: "cancelled".to_string(),
4810                mtp_drafted: 0,
4811                mtp_accepted: 0,
4812                token_confidence: Vec::new(),
4813                traces: Vec::new(),
4814            });
4815        }
4816
4817        // Prompt absorbed → freeze the o1 layers' skeletons; from here
4818        // every decode step on those layers is O(W + m·dv + m²).
4819        if !o1_sealed {
4820            match self.o1_seal_checked() {
4821                Ok(_) => {}
4822                Err(err) => {
4823                    self.finish_generation(&mut mtp, &mut router, true);
4824                    return Err(err);
4825                }
4826            }
4827        }
4828
4829        // Commit one token: push, check EOS, stream. Returns false = stop.
4830        macro_rules! commit {
4831            ($id:expr) => {{
4832                all_ids.push($id);
4833                generated += 1;
4834                self.note_draft_id($id);
4835                if self.tokenizer.is_eos($id) && !self.ignore_eos {
4836                    finish_reason = "stop".to_string();
4837                    false
4838                } else {
4839                    let token_text = self.tokenizer.decode_token($id);
4840                    let mut go = true;
4841                    if let Some(ref mut cb) = on_token {
4842                        if !cb(&token_text) {
4843                            finish_reason = "cancelled".to_string();
4844                            go = false;
4845                        }
4846                    }
4847                    go
4848                }
4849            }};
4850        }
4851
4852        // Speculation is decided by MEASUREMENT, not by an acceptance
4853        // model. A k=4 round costs ~3.8 plain tokens on the 5090 (draft
4854        // 6.6 + verify 66.6 + commit 4.8 ms against a 20.6 ms token), so it
4855        // pays only when the head lands ~2.8 of 4 — predictable text (code,
4856        // structured output) does, free prose often does not, and the
4857        // ratio at which the two cross depends on the card and the context
4858        // depth. So: four speculative rounds timed, then eight plain
4859        // tokens timed, and the faster arm runs until a re-check 256
4860        // tokens later (context growth moves the balance). The trial
4861        // costs at most a few tokens of the slower arm per 256.
4862        let mut spec_trial = SpecTrial::Spec {
4863            t0: std::time::Instant::now(),
4864            gen0: generated,
4865            rounds: 0,
4866        };
4867        // The token-count proxy prices a round at ~1.9 plain tokens. That
4868        // holds for the Metal rounds whose cost was measured — greedy and
4869        // the sparse sampling chain — so an expensive round (the dense
4870        // chain, reachable only by `CMF_GRAPH_SPEC_SAMPLE=1`) still times
4871        // the plain path before it decides.
4872        let mut spec_mon = SpecMon {
4873            metal: graph_spec && crate::gpu::q1_force() && spec_cheap_round,
4874            ..SpecMon::default()
4875        };
4876        let mut spec_watchdog_off = false;
4877        // CMF_GRAPH_SPEC_TIME: the round walls so far (round 1 excluded —
4878        // it pays the scratch), for the outlier test on each new one
4879        let mut spec_walls: Vec<f32> = Vec::new();
4880        // ... and the end of the last round: the host time between rounds
4881        // (token commits, streaming, the loop top) is printed at level 2
4882        let mut spec_round_end: Option<std::time::Instant> = None;
4883        if mimo_spec {
4884            if let Ok(path) = std::env::var("CMF_MIMO_MTP_PROBE") {
4885                if let Some(mut st) = self.mimo_mtp.take() {
4886                    self.mimo_mtp_probe(&mut st, input_ids, &path);
4887                    self.mimo_mtp = Some(st);
4888                }
4889            }
4890        }
4891        // ── Decode ──
4892        let mut next_pos = input_ids.len();
4893        'decode: while generated < max_tokens {
4894            if self
4895                .graph_failed
4896                .swap(false, std::sync::atomic::Ordering::Relaxed)
4897            {
4898                // Keep the detached MTP module attached after a terminal
4899                // graph error so the pipeline can be reused for a fresh
4900                // sequence.  `clear_sequence_state` only clears mirrors and
4901                // host KV; it cannot recover a module dropped here.
4902                self.finish_generation(&mut mtp, &mut router, true);
4903                return Err("GPU token graph failed during decode".to_string());
4904            }
4905            if self
4906                .cancel
4907                .swap(false, std::sync::atomic::Ordering::Relaxed)
4908            {
4909                finish_reason = "cancelled".to_string();
4910                break 'decode;
4911            }
4912            // A rejected speculative draft already drew this position's
4913            // token from the residual distribution (graph_spec_step); it
4914            // is committed as-is — sampling again from the row's logits
4915            // would bias the stream toward the target's mode.
4916            if mimo_spec && next_pos > 0 {
4917                // Every path leaves `hidden` = the backbone output at
4918                // next_pos-1; the draft layers read it (idempotent).
4919                self.mimo_note_rows(&hidden, next_pos - 1);
4920            }
4921            let forced = self.spec_forced.take();
4922            let mut logits = match (forced, self.graph_logits.take()) {
4923                (Some(_), _) => Vec::new(),
4924                (None, Some(lg)) => lg,
4925                (None, None) => {
4926                    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Head);
4927                    inference::rms_norm_into(
4928                        &hidden,
4929                        &self.weights.final_norm,
4930                        self.rms_eps,
4931                        self.norm_style,
4932                        &mut self.ws.n1,
4933                    );
4934                    self.lm_head_forward(&self.ws.n1)
4935                }
4936            };
4937            // CMF_LOGIT_DUMP=<path>: the first decode step's hidden + logits
4938            // as raw f32 (hidden first) — cross-backend numerics diffing.
4939            if generated
4940                == std::env::var("CMF_LOGIT_DUMP_STEP")
4941                    .ok()
4942                    .and_then(|v| v.parse().ok())
4943                    .unwrap_or(0)
4944            {
4945                if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
4946                    let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
4947                    for v in hidden.iter().chain(logits.iter()) {
4948                        bytes.extend_from_slice(&v.to_le_bytes());
4949                    }
4950                    if let Err(e) = std::fs::write(&path, &bytes) {
4951                        eprintln!("logit dump: failed to write {path}: {e}");
4952                        self.finish_generation(&mut mtp, &mut router, true);
4953                        return Err(format!("logit dump write failed: {e}"));
4954                    }
4955                }
4956            }
4957            // CMF_LOGIT_DUMP_ALL=<dir>: every decode step's logits as raw
4958            // f32, `<dir>/step{n:05}.f32` — step-by-step backend diffing
4959            // (a greedy run on two backends compares until they diverge).
4960            if let Ok(dir) = std::env::var("CMF_LOGIT_DUMP_ALL") {
4961                if !logits.is_empty() {
4962                    let path = std::path::Path::new(&dir).join(format!("step{generated:05}.f32"));
4963                    let bytes: Vec<u8> = logits.iter().flat_map(|v| v.to_le_bytes()).collect();
4964                    if let Err(e) =
4965                        std::fs::create_dir_all(&dir).and_then(|_| std::fs::write(&path, &bytes))
4966                    {
4967                        eprintln!("logit dump: failed to write {}: {e}", path.display());
4968                    }
4969                }
4970            }
4971            let t_next = match forced {
4972                Some(c) => c,
4973                None => {
4974                    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Sampler);
4975                    sampler::sample_with_scratch_pool(
4976                        &logits,
4977                        &self.sampler_config,
4978                        self.sampler_config.penalty_past(&all_ids, bounded_native),
4979                        &mut self.rng,
4980                        &mut self.sampler_scratch,
4981                        self.pool.as_deref(),
4982                    )
4983                }
4984            };
4985            if self.confidence_on {
4986                confidence.push(if logits.is_empty() {
4987                    0.0
4988                } else {
4989                    sampler::top1_prob_pool(
4990                        self.pool.as_deref(),
4991                        &mut self.sampler_scratch,
4992                        &logits,
4993                        t_next,
4994                        calib_temp,
4995                    )
4996                });
4997            }
4998            if !logits.is_empty() {
4999                attention::recycle_buf(&mut logits);
5000            }
5001            if trace_on {
5002                // active_skill = the overlay in force while this token was
5003                // generated; recon/switched are filled after the post-emit
5004                // routing eval below (freshest coherence for this token).
5005                let skill = router.as_ref().and_then(|r| r.active_id());
5006                traces.push(TokenTrace {
5007                    t: generated,
5008                    token_id: t_next,
5009                    confidence: confidence.last().copied().unwrap_or(0.0),
5010                    active_skill: skill,
5011                    recon: None,
5012                    switched: false,
5013                });
5014            }
5015            if !commit!(t_next) {
5016                break 'decode;
5017            }
5018            if generated >= max_tokens {
5019                break 'decode;
5020            }
5021
5022            if self.dsv41.is_none() && self.kv_cache.needs_eviction() {
5023                // Say it ONCE, loudly: past this point the model keeps
5024                // talking but has lost half its context, and on a GDN
5025                // hybrid the graph's device state goes stale on top. The
5026                // Qwen3.8 bring-up spent a day reading this cliff as
5027                // three different model bugs.
5028                static SAID: std::sync::Once = std::sync::Once::new();
5029                SAID.call_once(|| {
5030                    tracing::warn!(
5031                        "KV cache full at {} positions — evicting half; quality \
5032                         will degrade. Raise CMF_MAX_SEQ.",
5033                        self.kv_cache.max_seq_len,
5034                    );
5035                });
5036                let keep = (self.kv_cache.max_seq_len / 2).max(1);
5037                self.kv_cache.evict(keep);
5038            }
5039
5040            // Advance the speculation trial: plain-phase accounting and
5041            // the periodic re-check happen here, on every token.
5042            if graph_spec {
5043                match spec_trial {
5044                    SpecTrial::Plain { t0, gen0 } if spec_mon.plain_done(t0, gen0, generated) => {
5045                        spec_mon.plain_ms =
5046                            t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
5047                        let keep = spec_mon.pays();
5048                        tracing::info!(
5049                            "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5050                            spec_mon.tokens,
5051                            spec_mon.round_ms,
5052                            spec_mon.plain_ms,
5053                            if keep { "speculating" } else { "plain" }
5054                        );
5055                        spec_mon.fails = 0;
5056                        spec_trial = SpecTrial::Decided {
5057                            spec: keep,
5058                            recheck_at: if keep { usize::MAX } else { generated + 128 },
5059                        };
5060                    }
5061                    SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
5062                        spec_mon.n = 0;
5063                        spec_trial = SpecTrial::Spec {
5064                            t0: std::time::Instant::now(),
5065                            gen0: generated,
5066                            rounds: 0,
5067                        };
5068                    }
5069                    _ => {}
5070                }
5071                spec_watchdog_off = matches!(
5072                    spec_trial,
5073                    SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
5074                );
5075            }
5076            // ── MiMo-V2 draft stack: draft K, verify K+1 rows in one batch ──
5077            if mimo_spec && generated + 1 < max_tokens && next_pos > 0 {
5078                let budget = max_tokens - generated - 1;
5079                if let Some(mut st) = self.mimo_mtp.take() {
5080                    let k = st.depth.min(budget);
5081                    let r = self.mimo_spec_round(&mut st, next_pos, &all_ids, k);
5082                    self.mimo_mtp = Some(st);
5083                    let r = match r {
5084                        Ok(r) => r,
5085                        Err(err) => {
5086                            self.finish_generation(&mut mtp, &mut router, true);
5087                            return Err(err);
5088                        }
5089                    };
5090                    if let Some(r) = r {
5091                        drafted += r.drafted;
5092                        accepted += r.accepted.len();
5093                        let mut stopped = false;
5094                        for &id in &r.accepted {
5095                            if self.confidence_on {
5096                                confidence.push(0.0);
5097                            }
5098                            if !commit!(id) {
5099                                stopped = true;
5100                                break;
5101                            }
5102                        }
5103                        if stopped {
5104                            break 'decode;
5105                        }
5106                        next_pos += r.accepted.len() + 1;
5107                        hidden = r.hidden;
5108                        // The loop top chooses the round's own token from
5109                        // these logits — the same sampler, same history.
5110                        self.graph_logits = Some(r.logits);
5111                        continue 'decode;
5112                    }
5113                }
5114            }
5115            match &mut mtp {
5116                // ── Graph speculation: chain-draft, batch-verify on device ──
5117                #[cfg(feature = "gpu")]
5118                Some(m)
5119                    if graph_spec
5120                        && !spec_watchdog_off
5121                        && generated + 1 < max_tokens
5122                        && next_pos > 0 =>
5123                {
5124                    let t_round = std::time::Instant::now();
5125                    if spec_time_level() >= 2 {
5126                        if let Some(t) = spec_round_end.take() {
5127                            eprintln!(
5128                                "spec-gap {:.2} ms (host between rounds)",
5129                                t.elapsed().as_secs_f64() * 1e3
5130                            );
5131                        }
5132                    }
5133                    spec_stamps_begin();
5134                    // device buffers allocated during this round: a
5135                    // first-touch Shared allocation is zero-filled inside
5136                    // the command buffer that uses it, which is what the
5137                    // long outlier rounds were
5138                    #[cfg(target_os = "macos")]
5139                    let allocs0 = crate::gpu_metal::IO_BUF_ALLOCS
5140                        .load(std::sync::atomic::Ordering::Relaxed);
5141                    #[cfg(not(target_os = "macos"))]
5142                    let allocs0 = 0u64;
5143                    if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
5144                        m,
5145                        &hidden,
5146                        t_next,
5147                        next_pos,
5148                        &mut drafted,
5149                        &mut accepted,
5150                        &mut all_ids,
5151                        max_tokens - generated,
5152                    ) {
5153                        next_pos = n_pos;
5154                        hidden = new_h;
5155                        let level = spec_time_level();
5156                        if level > 0 {
5157                            let wall = t_round.elapsed().as_secs_f32() * 1e3;
5158                            let stamps = spec_stamps_take();
5159                            // the running median of the rounds before this
5160                            // one (round 1 pays the scratch: not a sample)
5161                            let median = if spec_walls.len() >= 3 {
5162                                let mut s = spec_walls.clone();
5163                                s.sort_by(|a, b| a.partial_cmp(b).unwrap());
5164                                Some(s[s.len() / 2])
5165                            } else {
5166                                None
5167                            };
5168                            let outlier = median.is_some_and(|m| wall > 1.4 * m);
5169                            #[cfg(target_os = "macos")]
5170                            let allocs = crate::gpu_metal::IO_BUF_ALLOCS
5171                                .load(std::sync::atomic::Ordering::Relaxed)
5172                                - allocs0;
5173                            #[cfg(not(target_os = "macos"))]
5174                            let allocs = allocs0;
5175                            eprintln!(
5176                                "spec-round wall {wall:.1} ms → {} tokens{}{}",
5177                                extra.len() + 1,
5178                                if allocs > 0 {
5179                                    format!(" [{allocs} new device buffers]")
5180                                } else {
5181                                    String::new()
5182                                },
5183                                match (outlier, median) {
5184                                    (true, Some(m)) => format!(" OUTLIER (median {m:.1})"),
5185                                    _ => String::new(),
5186                                }
5187                            );
5188                            if level >= 2 || outlier {
5189                                let sum: f32 = stamps.iter().map(|s| s.1).sum();
5190                                eprintln!(
5191                                    "spec-stamps: {}| untracked {:.1}",
5192                                    spec_stamps_format(&stamps),
5193                                    wall - sum
5194                                );
5195                            }
5196                            if spec_mon.n >= 1 {
5197                                spec_walls.push(wall);
5198                            }
5199                        }
5200                        // One speculative round done: the monitor counts it
5201                        // (round 1 untimed — it pays the batch scratch and
5202                        // the draft mirror), and the trial advances.
5203                        spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
5204                        // the round's tokens land in `generated` below; the
5205                        // plain phase must start counting AFTER them
5206                        spec_trial = Self::spec_trial_round(
5207                            spec_trial,
5208                            &mut spec_mon,
5209                            generated + extra.len() + 1,
5210                        );
5211                        let mut stopped = false;
5212                        for &id in &extra {
5213                            if self.confidence_on {
5214                                confidence.push(0.0);
5215                            }
5216                            if !commit!(id) {
5217                                stopped = true;
5218                                break;
5219                            }
5220                        }
5221                        if stopped {
5222                            break 'decode;
5223                        }
5224                        if spec_time_level() >= 2 {
5225                            spec_round_end = Some(std::time::Instant::now());
5226                        }
5227                        continue 'decode;
5228                    }
5229                    if self
5230                        .graph_failed
5231                        .swap(false, std::sync::atomic::Ordering::Relaxed)
5232                    {
5233                        // `graph_spec_step` may have detached MTP while a
5234                        // warm-up was in flight.  Do not reinterpret its
5235                        // terminal device failure as a plain decode step;
5236                        // restore the head, clear both mirrors, and surface
5237                        // one explicit error to the caller.
5238                        self.finish_generation(&mut mtp, &mut router, true);
5239                        return Err("GPU MTP graph failed during speculative decode".to_string());
5240                    }
5241                    // Declined (batch graph refused): plain forward below —
5242                    // and a round that produced one token for the trial's
5243                    // ledger, so a graph that keeps refusing is measured out
5244                    // like a head that keeps missing (it was spinning
5245                    // forever on a file whose batch graph declines).
5246                    // A declined round is not a cheap one-token round — it
5247                    // is a verify that does not exist for this file (a
5248                    // healed q8_2f tail measured 760 drafts, 0 accepted, 33
5249                    // against 48.8 tok/s while the monitor called the draft
5250                    // alone "paying"). Count it as the losing streak in one.
5251                    spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
5252                    spec_mon.tokens = 0.0;
5253                    spec_mon.fails = 3;
5254                    spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
5255                    hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
5256                    next_pos += 1;
5257                    continue 'decode;
5258                }
5259                // ── Speculative: draft t+2, verify in a fused pair ──
5260                Some(m) if !graph_spec && generated + 1 < max_tokens => {
5261                    let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
5262                    drafted += 1;
5263                    let emb1 = self.embed_single(t_next);
5264                    let emb2 = self.embed_single(draft);
5265                    let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
5266
5267                    inference::rms_norm_into(
5268                        &h1,
5269                        &self.weights.final_norm,
5270                        self.rms_eps,
5271                        self.norm_style,
5272                        &mut self.ws.n1,
5273                    );
5274                    let mut logits1 = self.lm_head_forward(&self.ws.n1);
5275                    let t_after = sampler::sample_with_scratch_pool(
5276                        &logits1,
5277                        &self.sampler_config,
5278                        self.sampler_config.penalty_past(&all_ids, bounded_native),
5279                        &mut self.rng,
5280                        &mut self.sampler_scratch,
5281                        self.pool.as_deref(),
5282                    );
5283                    if self.confidence_on {
5284                        confidence.push(sampler::top1_prob_pool(
5285                            self.pool.as_deref(),
5286                            &mut self.sampler_scratch,
5287                            &logits1,
5288                            t_after,
5289                            calib_temp,
5290                        ));
5291                    }
5292                    attention::recycle_buf(&mut logits1);
5293                    if trace_on {
5294                        // Speculative decode is mutually exclusive with
5295                        // dynamic routing (router is None here) — no skill.
5296                        traces.push(TokenTrace {
5297                            t: generated,
5298                            token_id: t_after,
5299                            confidence: confidence.last().copied().unwrap_or(0.0),
5300                            active_skill: None,
5301                            recon: None,
5302                            switched: false,
5303                        });
5304                    }
5305                    let stop = !commit!(t_after);
5306
5307                    if t_after == draft {
5308                        accepted += 1;
5309                        self.commit_linear_scratch();
5310                        let _ = self.mtp_step(m, &h1, t_after, next_pos);
5311                        hidden = h2;
5312                        next_pos += 2;
5313                    } else {
5314                        // The draft lane is wrong: roll its KV entry back.
5315                        for layer in &mut self.kv_cache.layers {
5316                            layer.truncate_last(1);
5317                        }
5318                        if !stop {
5319                            let _ = self.mtp_step(m, &h1, t_after, next_pos);
5320                            hidden = self.forward_layers(
5321                                &self.embed_single(t_after),
5322                                next_pos + 1,
5323                                None,
5324                            );
5325                        }
5326                        next_pos += 2;
5327                    }
5328                    if stop {
5329                        break 'decode;
5330                    }
5331                }
5332                // ── Vanilla: forward the sampled token ──
5333                _ => {
5334                    // ── DeepSeek-V4 speculative decode (CMF_DSV4_SPEC=1):
5335                    // draft five on the card, verify batched, commit the
5336                    // accepted prefix. Greedy only; a rejected token's state
5337                    // is restored and replayed, so output equals the walk. ──
5338                    #[cfg(feature = "gpu")]
5339                    if Self::dsv4_spec_on() && self.dsv4.is_some() {
5340                        static SAID: std::sync::Once = std::sync::Once::new();
5341                        SAID.call_once(|| {
5342                            eprintln!(
5343                                "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
5344                                !self.dsv4_mtp.is_empty(),
5345                                task_mask.is_none(),
5346                                router.is_none(),
5347                                !trace_on,
5348                                self.sampler_config.temperature < 1e-6,
5349                                self.sampler_config.repetition_penalty == 1.0,
5350                            );
5351                        });
5352                    }
5353                    #[cfg(feature = "gpu")]
5354                    if Self::dsv4_spec_on()
5355                        && self.dsv4.is_some()
5356                        && !self.dsv4_mtp.is_empty()
5357                        && task_mask.is_none()
5358                        && router.is_none()
5359                        && !trace_on
5360                        && self.sampler_config.temperature < 1e-6
5361                        && self.sampler_config.repetition_penalty == 1.0
5362                        && generated + 1 < max_tokens
5363                        && all_ids.len() >= 2
5364                        && generated >= dsv4_spec_retry_at
5365                    {
5366                        let tip_token = all_ids[all_ids.len() - 2];
5367                        let drafted0 = drafted;
5368                        let round = self.dsv4_spec_step(
5369                            tip_token,
5370                            t_next,
5371                            next_pos,
5372                            max_tokens.saturating_sub(generated),
5373                            &mut drafted,
5374                            &mut accepted,
5375                        );
5376                        if drafted > drafted0 {
5377                            let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
5378                            if useful {
5379                                dsv4_spec_bad = 0;
5380                            } else {
5381                                dsv4_spec_bad += 1;
5382                                if dsv4_spec_bad >= 2 {
5383                                    dsv4_spec_bad = 0;
5384                                    dsv4_spec_retry_at = generated.saturating_add(32);
5385                                    tracing::info!(
5386                                        "dsv4: draft не окупился дважды — точный walk на 32 токена"
5387                                    );
5388                                }
5389                            }
5390                        }
5391                        if let Some((extra, n_pos)) = round {
5392                            next_pos = n_pos;
5393                            let mut stopped = false;
5394                            for &id in &extra {
5395                                if self.confidence_on {
5396                                    confidence.push(0.0);
5397                                }
5398                                if !commit!(id) {
5399                                    stopped = true;
5400                                    break;
5401                                }
5402                            }
5403                            if stopped {
5404                                break 'decode;
5405                            }
5406                            continue 'decode;
5407                        }
5408                    }
5409                    self.graph_want_logits = fuse_lm;
5410                    // Greedy burst (CMF_MULTISTEP, default 8, 1 = off): while
5411                    // nothing observes per-token state — pure argmax sampling,
5412                    // no router/trace/confidence/mask — decode k tokens per
5413                    // submit and commit them wholesale. The trailing normal
5414                    // forward leaves logits for the loop top, as always.
5415                    let mut t_fwd = t_next;
5416                    let pure_greedy = self.sampler_config.temperature < 1e-6
5417                        && self.sampler_config.repetition_penalty == 1.0
5418                        && self.sampler_config.suppress_tokens.is_empty();
5419                    // Off by default: at every k the burst measured at or
5420                    // below the plain path on this graph shape (k=1 loses
5421                    // the argmax dispatches vs a 1 MB readback, k>=8 loses
5422                    // inter-step drains vs the saved sync). Experimental.
5423                    let burst_k = std::env::var("CMF_MULTISTEP")
5424                        .ok()
5425                        .and_then(|v| v.parse::<usize>().ok())
5426                        .unwrap_or(0);
5427                    if pure_greedy
5428                        && burst_k >= 1
5429                        && fuse_lm
5430                        && task_mask.is_none()
5431                        && router.is_none()
5432                        && !trace_on
5433                        && !self.confidence_on
5434                    {
5435                        let mut stopped = false;
5436                        loop {
5437                            let room = max_tokens.saturating_sub(generated);
5438                            if room <= 2 {
5439                                break;
5440                            }
5441                            let k = burst_k.min(room - 1);
5442                            if k < 1 {
5443                                break;
5444                            }
5445                            let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
5446                                if self
5447                                    .graph_failed
5448                                    .swap(false, std::sync::atomic::Ordering::Relaxed)
5449                                {
5450                                    self.finish_generation(&mut mtp, &mut router, true);
5451                                    return Err(
5452                                        "GPU token graph failed during greedy burst".to_string()
5453                                    );
5454                                }
5455                                break;
5456                            };
5457                            next_pos += k;
5458                            for &id in &ids {
5459                                if !commit!(id) {
5460                                    stopped = true;
5461                                    break;
5462                                }
5463                            }
5464                            if stopped {
5465                                break;
5466                            }
5467                            t_fwd = *ids.last().unwrap();
5468                        }
5469                        if stopped {
5470                            break 'decode;
5471                        }
5472                    }
5473                    // Metal: keep the draft head's cache in step through
5474                    // the trial's plain phase and a paused speculation —
5475                    // the pair (hidden, t_fwd) at next_pos−1, the step the
5476                    // round's draft 0 would take. Without it the head's
5477                    // cache lagged the trunk by every plain token for the
5478                    // rest of the generation: the batched warm-up declined
5479                    // every later round and its rows went one by one (a
5480                    // whole MTP step per accepted token), and the drafts
5481                    // attended a context with those tokens missing.
5482                    #[cfg(target_os = "macos")]
5483                    if graph_spec
5484                        && spec_watchdog_off
5485                        && next_pos > 0
5486                        && self.mtp_graph_mode == Some(true)
5487                        && crate::gpu::q1_force()
5488                    {
5489                        if let Some(m) = mtp.as_mut() {
5490                            let _ = self.mtp_step_metal(m, &hidden, t_fwd, next_pos - 1, false);
5491                        }
5492                    }
5493                    hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
5494                    next_pos += 1;
5495                    // Dynamic routing: the forward updated φ; ask the
5496                    // router whether to switch skills before the next token.
5497                    if let Some(r) = &mut router {
5498                        let phi = self.dyn_phi_ema.clone();
5499                        let decision = r.step(&phi, generated);
5500                        if let Some(new_active) = decision {
5501                            let _ = self.set_active_skill(new_active);
5502                        }
5503                        // Backfill this token's coherence + switch flag from
5504                        // the just-run eval (freshest measured values).
5505                        if trace_on {
5506                            if let Some(last) = traces.last_mut() {
5507                                let e = r.last_best_e();
5508                                last.recon = e.is_finite().then_some(e);
5509                                last.switched = decision.is_some();
5510                            }
5511                        }
5512                    }
5513                }
5514            }
5515        }
5516
5517        let cancelled = finish_reason == "cancelled";
5518        // A generation during which the router switched weights holds no
5519        // state any single overlay would produce (each switch cleared the
5520        // cache mid-sequence), so it leaves no reuse key behind either.
5521        let dyn_switched = router.as_ref().is_some_and(|r| !r.switches.is_empty());
5522        if mimo_spec {
5523            if let Some(st) = self.mimo_mtp.as_ref() {
5524                let line = st.stats.line();
5525                tracing::info!("{line}");
5526                if std::env::var_os("CMF_MIMO_MTP_STATS").is_some() {
5527                    eprintln!("{line}");
5528                }
5529            }
5530        }
5531        self.finish_generation(&mut mtp, &mut router, cancelled);
5532
5533        let output_ids = &all_ids[input_ids.len()..];
5534        // Forwarded = prompt + all generated but the LAST sampled token
5535        // (emitted without being fed back). Exact only without MTP —
5536        // reuse is gated off when MTP is active.
5537        let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
5538        // A MiMo speculative round that stopped on an accepted draft (EOS,
5539        // cancel) leaves verify rows past the committed stream in the cache:
5540        // never offer that cache for reuse. Neither does a router that
5541        // switched weights mid-sequence (no single overlay produced it).
5542        let consumed = std::mem::take(&mut all_ids);
5543        if dyn_switched {
5544            self.clear_sequence_state();
5545        } else if cancelled || mimo_spec || prompt_rows.is_some() {
5546            self.clear_history();
5547        } else {
5548            self.record_consumed_prefix(&consumed[..forwarded.min(consumed.len())], reuse_from);
5549        }
5550        all_ids = consumed;
5551        let output_ids = &all_ids[input_ids.len()..];
5552        confidence.truncate(output_ids.len()); // guard against any overshoot
5553        traces.truncate(output_ids.len());
5554        Ok(GenerateResult {
5555            text: self.tokenizer.decode(output_ids),
5556            token_ids: output_ids.to_vec(),
5557            prompt_tokens: input_ids.len(),
5558            tokens_generated: generated,
5559            finish_reason,
5560            mtp_drafted: drafted,
5561            mtp_accepted: accepted,
5562            token_confidence: confidence,
5563            traces,
5564        })
5565    }
5566
5567    /// One MTP step: feed `(hidden_p, token_{p+1})` into the draft head,
5568    /// advance its KV cache at position `p`, return the drafted token
5569    /// for position `p+2`.
5570    fn mtp_step(
5571        &mut self,
5572        m: &mut MtpModule,
5573        hidden: &[f32],
5574        next_token: u32,
5575        position: usize,
5576    ) -> u32 {
5577        self.mtp_step_h(m, hidden, next_token, position).0
5578    }
5579
5580    /// Tally for `CMF_MTP_CHAIN_PROBE`: per depth, how often the CHAIN is
5581    /// still an exact prefix of the real continuation. Printed every 128
5582    /// depth-0 samples so a killed run still shows its table.
5583    fn chain_probe_note(depth: usize, prefix_ok: bool) {
5584        use std::sync::Mutex;
5585        static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
5586        let mut t = T.lock().unwrap();
5587        if t.len() <= depth {
5588            t.resize(depth + 1, (0, 0));
5589        }
5590        t[depth].0 += 1;
5591        t[depth].1 += prefix_ok as u64;
5592        if depth == 0 && t[0].0 % 128 == 0 {
5593            let line: Vec<String> = t
5594                .iter()
5595                .enumerate()
5596                .map(|(d, (n, k))| {
5597                    format!(
5598                        "d{}={:.0}%({n})",
5599                        d + 1,
5600                        100.0 * *k as f64 / (*n).max(1) as f64
5601                    )
5602                })
5603                .collect();
5604            eprintln!("mtp-chain: {}", line.join(" "));
5605        }
5606    }
5607
5608    /// `mtp_step` that also hands back the block's own output hidden — the
5609    /// state a CHAINED draft feeds the next step, the way a multi-token
5610    /// speculative round iterates the head on itself.
5611    /// One MTP block step from (trunk hidden, token): the head's LOGITS
5612    /// and the block's own hidden for chaining. The draft is argmax of the
5613    /// logits on the greedy path and a draw from their post-chain
5614    /// distribution on the sampling path.
5615    fn mtp_step_hl(
5616        &mut self,
5617        m: &mut MtpModule,
5618        hidden: &[f32],
5619        next_token: u32,
5620        position: usize,
5621    ) -> (Vec<f32>, Vec<f32>) {
5622        // The graph arm: the MTP block as a one-layer token graph with the
5623        // head fused — device attention over the block's own KV mirror,
5624        // one submit for block + head, hidden and logits back together.
5625        // Decided once per generation (see `mtp_graph_mode`).
5626        #[cfg(target_os = "macos")]
5627        if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
5628            if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
5629                self.mtp_graph_mode = Some(true);
5630                return r;
5631            }
5632            if self.mtp_graph_mode == Some(true) {
5633                tracing::error!("mtp Metal graph failed after admission");
5634                self.clear_sequence_state();
5635                self.graph_failed
5636                    .store(true, std::sync::atomic::Ordering::Relaxed);
5637                self.cancel
5638                    .store(true, std::sync::atomic::Ordering::Relaxed);
5639                return (Vec::new(), Vec::new());
5640            }
5641            self.mtp_graph_mode = Some(false);
5642        }
5643        #[cfg(feature = "gpu")]
5644        if self.mtp_graph_mode != Some(false) {
5645            if !self.mtp_graph_ok(m) {
5646                if self.mtp_graph_mode == Some(true) {
5647                    // A mirror was already admitted, so a capability change
5648                    // cannot safely switch this request to the stale CPU
5649                    // cache.  Keep the same terminal contract as a failed
5650                    // token graph.
5651                    tracing::error!("mtp graph became unavailable after admission");
5652                    self.clear_sequence_state();
5653                    self.graph_failed
5654                        .store(true, std::sync::atomic::Ordering::Relaxed);
5655                    self.cancel
5656                        .store(true, std::sync::atomic::Ordering::Relaxed);
5657                    return (Vec::new(), Vec::new());
5658                }
5659                self.mtp_graph_mode = Some(false);
5660            } else {
5661                if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
5662                    self.mtp_graph_mode = Some(true);
5663                    return r;
5664                }
5665                if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
5666                    // A token graph can have admitted a persistent MTP/GDN
5667                    // mirror before its readback failed.  The CPU MTP cache
5668                    // is not a valid continuation in that state; leave the
5669                    // flag set so the generation caller returns through its
5670                    // terminal error path instead of silently switching
5671                    // arithmetic.
5672                    return (Vec::new(), Vec::new());
5673                }
5674                // `mtp_graph_ok` was true, so a None here means a refusal or
5675                // failure after graph admission.  Do not fall through to a
5676                // CPU cache whose rows may lag the device mirror.
5677                tracing::error!("mtp graph failed or declined after admission");
5678                self.clear_sequence_state();
5679                self.graph_failed
5680                    .store(true, std::sync::atomic::Ordering::Relaxed);
5681                self.cancel
5682                    .store(true, std::sync::atomic::Ordering::Relaxed);
5683                return (Vec::new(), Vec::new());
5684            }
5685        }
5686        // fc concat order is [enorm(embed); hnorm(hidden)] — EMBEDDING
5687        // FIRST. Verified by the oracle (converter/mtp_oracle.py):
5688        // [emb;hid] → 45.8% acceptance, [hid;emb] → 0.00%.
5689        let e = self.embed_single(next_token);
5690        let mut cat = vec![0.0f32; 2 * self.hidden_size];
5691        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5692        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5693        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5694        let mut x = vec![0.0f32; self.hidden_size];
5695        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5696
5697        // One standard transformer block over the MTP's own cache.
5698        let lw = &m.layer;
5699        inference::rms_norm_into(
5700            &x,
5701            &lw.input_norm,
5702            self.rms_eps,
5703            self.norm_style,
5704            &mut self.ws.n1,
5705        );
5706        let attn = match &lw.attn {
5707            // MLA models carry no MTP head; this path cannot see them.
5708            AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
5709            AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
5710            AttnKind::Full {
5711                wq,
5712                wk,
5713                wv,
5714                wo,
5715                q_norm,
5716                k_norm,
5717                output_gate,
5718                softplus_gate,
5719                bias,
5720            } => {
5721                let mut cfg = self.attn_cfg(position);
5722                cfg.q_norm = q_norm.as_deref();
5723                cfg.k_norm = k_norm.as_deref();
5724                cfg.output_gate = *output_gate;
5725                cfg.softplus_gate = softplus_gate
5726                    .as_ref()
5727                    .map(|(gate, per_head)| (gate, *per_head));
5728                cfg.bias = bias
5729                    .as_ref()
5730                    .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
5731                attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
5732            }
5733            AttnKind::Linear(_)
5734            | AttnKind::LinearGdn(_)
5735            | AttnKind::ShortConv(_)
5736            | AttnKind::Bounded(_) => {
5737                unreachable!("MTP block is full attention")
5738            }
5739        };
5740        for (i, &a) in attn.iter().enumerate() {
5741            x[i] += a;
5742        }
5743        inference::rms_norm_into(
5744            &x,
5745            &lw.post_norm,
5746            self.rms_eps,
5747            self.norm_style,
5748            &mut self.ws.p1,
5749        );
5750        let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
5751        for (i, &f) in ffn.iter().enumerate() {
5752            x[i] += f;
5753        }
5754
5755        inference::rms_norm_into(
5756            &x,
5757            &m.final_norm,
5758            self.rms_eps,
5759            self.norm_style,
5760            &mut self.ws.n1,
5761        );
5762        let lg = self.lm_head_forward(&self.ws.n1);
5763        (lg, x)
5764    }
5765
5766    /// `mtp_step_hl` reduced to the greedy draft: argmax of the head.
5767    fn mtp_step_h(
5768        &mut self,
5769        m: &mut MtpModule,
5770        hidden: &[f32],
5771        next_token: u32,
5772        position: usize,
5773    ) -> (u32, Vec<f32>) {
5774        let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
5775        let draft = sampler::argmax(&lg);
5776        attention::recycle_buf(&mut lg);
5777        (draft, x)
5778    }
5779
5780    /// One speculative round for the trial: rounds 1..5 of a `Spec` phase
5781    /// advance it (the monitor already averaged this round); after five,
5782    /// the plain phase runs (once — a known plain rate decides at once);
5783    /// a decided speculation keeps re-checking the rule every round and
5784    /// stops after four losing rounds in a row.
5785    fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
5786        match trial {
5787            SpecTrial::Spec { t0, gen0, rounds } => {
5788                let rounds = rounds + 1;
5789                if rounds >= 5 {
5790                    if mon.plain_ms > 0.0 {
5791                        let keep = mon.pays();
5792                        mon.fails = 0;
5793                        tracing::info!(
5794                            "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5795                            mon.tokens,
5796                            mon.round_ms,
5797                            mon.plain_ms,
5798                            if keep { "speculating" } else { "plain" }
5799                        );
5800                        SpecTrial::Decided {
5801                            spec: keep,
5802                            recheck_at: if keep { usize::MAX } else { generated + 128 },
5803                        }
5804                    } else if mon.pays() {
5805                        // Metal: the rounds land enough tokens each that no
5806                        // plain measurement is needed — keep speculating,
5807                        // and re-check every round (a losing streak sends
5808                        // the loop to the plain phase, below).
5809                        mon.fails = 0;
5810                        tracing::info!(
5811                            "speculation trial: {:.2} tok/round in {:.1} ms — speculating (plain not timed)",
5812                            mon.tokens,
5813                            mon.round_ms,
5814                        );
5815                        SpecTrial::Decided {
5816                            spec: true,
5817                            recheck_at: usize::MAX,
5818                        }
5819                    } else {
5820                        SpecTrial::Plain {
5821                            t0: std::time::Instant::now(),
5822                            gen0: generated,
5823                        }
5824                    }
5825                } else {
5826                    SpecTrial::Spec { t0, gen0, rounds }
5827                }
5828            }
5829            SpecTrial::Decided { spec: true, .. } => {
5830                if mon.pays() {
5831                    mon.fails = 0;
5832                    trial
5833                } else {
5834                    mon.fails += 1;
5835                    if mon.fails >= 4 {
5836                        if mon.plain_ms <= 0.0 {
5837                            // Metal, plain never timed: four doubtful rounds
5838                            // buy the (bounded) plain measurement, and the
5839                            // exact rule decides from it.
5840                            tracing::info!(
5841                                "speculation doubtful: {:.2} tok/round in {:.1} ms — timing plain",
5842                                mon.tokens,
5843                                mon.round_ms,
5844                            );
5845                            return SpecTrial::Plain {
5846                                t0: std::time::Instant::now(),
5847                                gen0: generated,
5848                            };
5849                        }
5850                        tracing::info!(
5851                            "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
5852                            mon.tokens,
5853                            mon.round_ms,
5854                            mon.plain_ms
5855                        );
5856                        SpecTrial::Decided {
5857                            spec: false,
5858                            recheck_at: generated + 128,
5859                        }
5860                    } else {
5861                        trial
5862                    }
5863                }
5864            }
5865            other => other,
5866        }
5867    }
5868
5869    /// The MTP block's device-mirror id: the trunk's id with a high bit,
5870    /// so the (kv_id, layer) mirror keys never collide.
5871    fn mtp_kv_id(&self) -> u64 {
5872        self.graph_kv_id | (1u64 << 40)
5873    }
5874
5875    /// The MTP block's mirror layer index: 0 — its own kv_id keeps it
5876    /// apart from the trunk, and the BATCH graph (the warm-up path) keys
5877    /// its mirrors at layer 0 with no base of its own, so the draft's
5878    /// token graph must key the same slot.
5879    const MTP_LAYER_BASE: usize = 0;
5880
5881    /// The wgpu MTP draft writes speculative rows straight into its device
5882    /// mirror while the CPU owner retains only the real prompt/decode anchor.
5883    /// After verification, move that mirror cursor back to the anchor before
5884    /// replaying accepted pairs.  The next graph append then sees the same
5885    /// contiguous position as the CPU/Metal path without uploading stale
5886    /// speculative rows.
5887    #[cfg(feature = "gpu")]
5888    fn rewind_mtp_graph_mirror(&self, stored: usize) -> bool {
5889        self.mtp_graph_mode != Some(true)
5890            || crate::gpu::graph_kv_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, stored)
5891    }
5892
5893    /// A speculative verify graph appends the full `k+1` trunk rows before
5894    /// the acceptance count is known.  GDN state already has a snapshot
5895    /// restore; Full-attention mirrors need the matching logical cursor
5896    /// rewind so the next graph call does not reject an ahead-of-position KV
5897    /// cache after a partial acceptance.
5898    #[cfg(feature = "gpu")]
5899    fn rewind_trunk_graph_mirrors(&self, stored: usize) -> bool {
5900        let mut ok = true;
5901        let mut expected = false;
5902        for li in 0..self.num_layers {
5903            if matches!(
5904                self.weights.layers[self.phys_layer(li)].attn,
5905                AttnKind::Full { .. }
5906            ) {
5907                expected = true;
5908                ok &= crate::gpu::graph_kv_set_stored(self.graph_kv_id, li, stored);
5909            }
5910        }
5911        !expected || ok
5912    }
5913
5914    /// Count the recurrent layers participating in the trunk verify graph.
5915    /// Snapshot restore is all-or-nothing across that set; deriving the count
5916    /// from the model keeps the restore contract valid for looped models too.
5917    fn graph_gdn_layer_count(&self) -> usize {
5918        (0..self.num_layers)
5919            .filter(|&li| {
5920                matches!(
5921                    &self.weights.layers[self.phys_layer(li)].attn,
5922                    AttnKind::LinearGdn(_)
5923                )
5924            })
5925            .count()
5926    }
5927
5928    /// The block's input from (trunk hidden, token): eh_proj · [enorm(e);
5929    /// hnorm(h)] — the same arithmetic the per-op path starts with.
5930    fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
5931        let e = self.embed_single(next_token);
5932        let mut cat = vec![0.0f32; 2 * self.hidden_size];
5933        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5934        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5935        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5936        let mut x = vec![0.0f32; self.hidden_size];
5937        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5938        x
5939    }
5940
5941    /// Is the MTP block graphable at all (device up, full attention
5942    /// without softplus, dense FFN)? The plan itself is built per call.
5943    #[cfg(feature = "gpu")]
5944    fn mtp_block_graph_ok(&self, m: &MtpModule) -> bool {
5945        if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
5946            return false;
5947        }
5948        if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
5949            || !crate::gpu::enabled_here()
5950            || self.attn_softcap > 0.0
5951            || self.attention_heads_per_layer.is_some()
5952            // The block graph caches V as wide as K and feeds o_proj
5953            // nh·head_dim; a narrow-V model keeps its MTP block per-op.
5954            || self.v_head_dim.is_some()
5955        {
5956            return false;
5957        }
5958        matches!(
5959            &m.layer.attn,
5960            AttnKind::Full {
5961                softplus_gate: None,
5962                ..
5963            }
5964        ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
5965    }
5966
5967    /// Full MTP token-graph eligibility, including the fused lm-head and all
5968    /// block projection weights.  Keep this distinct from the block-only
5969    /// check: prompt warm-up does not need the head, while a draft step does.
5970    #[cfg(feature = "gpu")]
5971    fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
5972        if !self.mtp_block_graph_ok(m) {
5973            return false;
5974        }
5975        let AttnKind::Full { wq, wk, wv, wo, .. } = &m.layer.attn else {
5976            return false;
5977        };
5978        let FfnKind::Dense(d) = &m.layer.ffn else {
5979            return false;
5980        };
5981        d.segs.is_empty()
5982            && wq.graph_weight().is_some()
5983            && wk.graph_weight().is_some()
5984            && wv.graph_weight().is_some()
5985            && wo.graph_weight().is_some()
5986            && d.gate_proj.graph_weight().is_some()
5987            && d.up_proj.graph_weight().is_some()
5988            && d.down_proj.graph_weight().is_some()
5989            && self.weights.lm_head.graph_weight().is_some()
5990    }
5991
5992    /// One MTP block step on the wgpu token graph: block + fused head in
5993    /// one submit, the block hidden and the logits read back together.
5994    /// None = the graph cannot take this block (softplus gate, non-dense
5995    /// FFN, unquantized head, no device) — the caller keeps the per-op
5996    /// path for the whole generation.
5997    #[cfg(feature = "gpu")]
5998    fn mtp_step_graph(
5999        &mut self,
6000        m: &mut MtpModule,
6001        hidden: &[f32],
6002        next_token: u32,
6003        position: usize,
6004    ) -> Option<(Vec<f32>, Vec<f32>)> {
6005        if !self.mtp_graph_ok(m) {
6006            return None;
6007        }
6008        let lw = &m.layer;
6009        let AttnKind::Full {
6010            wq,
6011            wk,
6012            wv,
6013            wo,
6014            q_norm,
6015            k_norm,
6016            output_gate,
6017            softplus_gate,
6018            bias,
6019        } = &lw.attn
6020        else {
6021            return None;
6022        };
6023        if softplus_gate.is_some() {
6024            return None;
6025        }
6026        let FfnKind::Dense(d) = &lw.ffn else {
6027            return None;
6028        };
6029        if !d.segs.is_empty() {
6030            return None; // tube layers run on the segmented path
6031        }
6032        // The block's input first: it borrows `self` mutably (embed scratch,
6033        // pool), the plan below borrows the weights immutably.
6034        let mut x = self.mtp_block_input(m, hidden, next_token);
6035        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6036            let (_, i, kind, rs) = t.graph_weight()?;
6037            Some(crate::gpu::GraphW {
6038                idx: i,
6039                kind,
6040                row_scale: rs,
6041                data: &[],
6042                prism: crate::gpu::GraphPrismOp::None,
6043                affine: false,
6044            })
6045        }
6046        let (model, _, _, _) = wq.graph_weight()?;
6047        let model = model.clone();
6048        let (lm_gw, lm_rows) = {
6049            let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6050            // The draft's head over the CMF_DRAFT_VOCAB shortlist (the same
6051            // cut the native Metal draft takes): 662 MB a step on Qwen3.8
6052            // becomes 170 MB at 65536; the verify keeps the full head.
6053            let rows = if kind == 6 {
6054                self.draft_head_rows(self.weights.lm_head.rows())
6055            } else {
6056                self.weights.lm_head.rows()
6057            };
6058            (
6059                crate::gpu::GraphW {
6060                    idx: i,
6061                    kind,
6062                    row_scale: rs,
6063                    data: &[],
6064                    prism: crate::gpu::GraphPrismOp::None,
6065                    affine: false,
6066                },
6067                rows,
6068            )
6069        };
6070        let layer = crate::gpu::GraphLayer {
6071            input_norm: &lw.input_norm,
6072            attn: crate::gpu::GraphAttn::Full {
6073                wq: gw(wq)?,
6074                wk: gw(wk)?,
6075                wv: gw(wv)?,
6076                wo: gw(wo)?,
6077                q_norm: q_norm.as_deref(),
6078                k_norm: k_norm.as_deref(),
6079                late_qk_norm: self.qk_norm_after_rope,
6080                bias: bias
6081                    .as_ref()
6082                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6083                output_gate: *output_gate,
6084                cpu_k: m.kv.k_heads(),
6085                cpu_v: m.kv.v_heads(),
6086                geom: None,
6087            },
6088            post_norm: &lw.post_norm,
6089            ffn: crate::gpu::GraphFfn::Dense {
6090                gate: gw(&d.gate_proj)?,
6091                up: gw(&d.up_proj)?,
6092                down: gw(&d.down_proj)?,
6093            },
6094        };
6095        let nh = self.num_heads;
6096        let (nkv, hd, rd) = self.layer_geom(0);
6097        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6098        let mut logits = Vec::new();
6099        let ok = crate::gpu::forward_token_graph(
6100            &model,
6101            self.mtp_kv_id(),
6102            std::slice::from_ref(&layer),
6103            &[None],
6104            self.o1_epoch,
6105            &self.inv_freq,
6106            &mut x,
6107            nh,
6108            nkv,
6109            hd,
6110            self.attn_scale,
6111            rd,
6112            self.hidden_size,
6113            self.intermediate_size,
6114            position,
6115            self.kv_cache.max_seq_len,
6116            gemma,
6117            self.rms_eps as f32,
6118            Some((&lm_gw, lm_rows)),
6119            &m.final_norm,
6120            &mut logits,
6121            &[],
6122            1,
6123            None,
6124            None,
6125            None,
6126            Self::MTP_LAYER_BASE,
6127            true,
6128        );
6129        match ok {
6130            crate::gpu::TokenGraphOutcome::Completed => {}
6131            crate::gpu::TokenGraphOutcome::Declined => return None,
6132            crate::gpu::TokenGraphOutcome::Failed => {
6133                // The backend has already admitted persistent state.  Keep
6134                // this distinct from a capability refusal so the caller
6135                // cannot switch to the stale CPU MTP cache.
6136                self.clear_sequence_state();
6137                self.graph_failed
6138                    .store(true, std::sync::atomic::Ordering::Relaxed);
6139                self.cancel
6140                    .store(true, std::sync::atomic::Ordering::Relaxed);
6141                return None;
6142            }
6143        }
6144        logits.resize(self.vocab_size, 0.0);
6145        Some((logits, x))
6146    }
6147
6148    /// The warm-ups of one speculative round on the device: every accepted
6149    /// (hidden, token) pair as ONE batched graph run over the MTP block
6150    /// (no head) — its kv_append lands the pairs in the block's mirror.
6151    /// `pairs` are consecutive positions from `first_pos`.  The tri-state
6152    /// result is intentional: a refusal before admission may use the
6153    /// per-row/CPU route, while a failure after admission must terminate the
6154    /// sequence rather than fall through to a stale CPU cache.
6155    #[cfg(feature = "gpu")]
6156    fn mtp_warm_graph(
6157        &mut self,
6158        m: &mut MtpModule,
6159        pairs: &[(&[f32], u32)],
6160        first_pos: usize,
6161    ) -> crate::gpu::BatchGraphOutcome {
6162        if pairs.is_empty() {
6163            return crate::gpu::BatchGraphOutcome::Completed;
6164        }
6165        if !self.mtp_block_graph_ok(m) {
6166            return crate::gpu::BatchGraphOutcome::Declined;
6167        }
6168        let hs = self.hidden_size;
6169        // Block inputs for every pair (eh_proj on the per-op path, one
6170        // matvec each — the plan's own prologue).
6171        let mut hiddens = Vec::with_capacity(pairs.len() * hs);
6172        for (h, t) in pairs {
6173            hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
6174        }
6175        let lw = &m.layer;
6176        let AttnKind::Full {
6177            wq,
6178            wk,
6179            wv,
6180            wo,
6181            q_norm,
6182            k_norm,
6183            output_gate,
6184            bias,
6185            ..
6186        } = &lw.attn
6187        else {
6188            return crate::gpu::BatchGraphOutcome::Declined;
6189        };
6190        let FfnKind::Dense(d) = &lw.ffn else {
6191            return crate::gpu::BatchGraphOutcome::Declined;
6192        };
6193        if !d.segs.is_empty() {
6194            return crate::gpu::BatchGraphOutcome::Declined; // tube layers run on the segmented path
6195        }
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 Some((model, _, _, _)) = wq.graph_weight() else {
6208            return crate::gpu::BatchGraphOutcome::Declined;
6209        };
6210        let model = model.clone();
6211        let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
6212            gw(wq),
6213            gw(wk),
6214            gw(wv),
6215            gw(wo),
6216            gw(&d.gate_proj),
6217            gw(&d.up_proj),
6218            gw(&d.down_proj),
6219        ) else {
6220            return crate::gpu::BatchGraphOutcome::Declined;
6221        };
6222        let layer = crate::gpu::GraphLayer {
6223            input_norm: &lw.input_norm,
6224            attn: crate::gpu::GraphAttn::Full {
6225                wq: gwq,
6226                wk: gwk,
6227                wv: gwv,
6228                wo: gwo,
6229                q_norm: q_norm.as_deref(),
6230                k_norm: k_norm.as_deref(),
6231                late_qk_norm: self.qk_norm_after_rope,
6232                bias: bias
6233                    .as_ref()
6234                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6235                output_gate: *output_gate,
6236                cpu_k: m.kv.k_heads(),
6237                cpu_v: m.kv.v_heads(),
6238                geom: None,
6239            },
6240            post_norm: &lw.post_norm,
6241            ffn: crate::gpu::GraphFfn::Dense {
6242                gate: gg,
6243                up: gu,
6244                down: gd,
6245            },
6246        };
6247        let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
6248        let nh = self.num_heads;
6249        let (nkv, hd, rd) = self.layer_geom(0);
6250        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6251        crate::gpu::forward_batch_graph(
6252            &model,
6253            self.mtp_kv_id(),
6254            std::slice::from_ref(&layer),
6255            &self.inv_freq,
6256            &mut hiddens,
6257            nh,
6258            nkv,
6259            hd,
6260            rd,
6261            hs,
6262            self.intermediate_size,
6263            &positions,
6264            self.kv_cache.max_seq_len,
6265            gemma,
6266            self.rms_eps as f32,
6267            self.attn_scale,
6268            pairs.len(),
6269            &[],
6270            0,
6271            None,
6272            None,
6273        )
6274    }
6275
6276    /// Complete an MTP warm-up after the batched graph has refused.  A
6277    /// graphable block is retried one row at a time; once any device row has
6278    /// been admitted, a CPU fallback would observe a stale mirror, so every
6279    /// token-graph refusal is terminal.  If the block is not graphable and no
6280    /// mirror exists yet, warming on the CPU is safe and records the CPU mode
6281    /// for the rest of the generation.
6282    #[cfg(feature = "gpu")]
6283    fn mtp_warm_graph_fallback(
6284        &mut self,
6285        m: &mut MtpModule,
6286        pairs: &[(&[f32], u32)],
6287        first_pos: usize,
6288    ) -> bool {
6289        if pairs.is_empty() {
6290            return true;
6291        }
6292        let graphable = self.mtp_block_graph_ok(m);
6293        if !graphable {
6294            // A previously admitted mirror cannot be made coherent by
6295            // appending to the host cache.  The caller turns this into a
6296            // terminal generation error and clears both mirrors.
6297            if self.mtp_graph_mode == Some(true) {
6298                return false;
6299            }
6300            self.mtp_graph_mode = Some(false);
6301            for (j, (h, t)) in pairs.iter().enumerate() {
6302                self.mtp_warm(m, h, *t, first_pos + j);
6303            }
6304            return true;
6305        }
6306
6307        // The batch refusal is recoverable only through the same device
6308        // state.  Keep rows owned until each token graph has completed; a
6309        // None is treated as unsafe because the token-graph API deliberately
6310        // collapses its backend refusal/failure into that result.
6311        for (j, (h, t)) in pairs.iter().enumerate() {
6312            if self.mtp_step_graph(m, h, *t, first_pos + j).is_none() {
6313                return false;
6314            }
6315        }
6316        self.mtp_graph_mode = Some(true);
6317        true
6318    }
6319
6320    /// Warm a contiguous set of MTP pairs using the existing graph seam, with
6321    /// an all-or-nothing error contract for callers that already admitted the
6322    /// trunk batch.  The non-GPU build keeps the same pair accounting while
6323    /// using the established CPU warm path.
6324    #[cfg(feature = "gpu")]
6325    fn mtp_warm_prefill_pairs(
6326        &mut self,
6327        m: &mut MtpModule,
6328        pairs: &[(&[f32], u32)],
6329        first_pos: usize,
6330    ) -> Result<(), &'static str> {
6331        // Keep unsupported token-graph heads on the established CPU MTP
6332        // route before admitting any block mirror.  Once a device mirror is
6333        // active, the same condition is terminal because CPU rows cannot
6334        // repair its state.
6335        if self.mtp_graph_mode == Some(false) || !self.mtp_graph_ok(m) {
6336            if self.mtp_graph_mode == Some(true) {
6337                return Err("MTP token graph became unavailable after admission");
6338            }
6339            self.mtp_graph_mode = Some(false);
6340            for (j, (h, t)) in pairs.iter().enumerate() {
6341                self.mtp_warm(m, h, *t, first_pos + j);
6342            }
6343            return Ok(());
6344        }
6345        match self.mtp_warm_graph(m, pairs, first_pos) {
6346            crate::gpu::BatchGraphOutcome::Completed => {
6347                if !pairs.is_empty() {
6348                    self.mtp_graph_mode = Some(true);
6349                }
6350                Ok(())
6351            }
6352            crate::gpu::BatchGraphOutcome::Declined => {
6353                if self.mtp_warm_graph_fallback(m, pairs, first_pos) {
6354                    Ok(())
6355                } else {
6356                    Err("MTP warm-up fallback failed after device admission")
6357                }
6358            }
6359            crate::gpu::BatchGraphOutcome::Failed => {
6360                Err("MTP warm batch graph failed after admission")
6361            }
6362        }
6363    }
6364
6365    #[cfg(not(feature = "gpu"))]
6366    fn mtp_warm_prefill_pairs(
6367        &mut self,
6368        m: &mut MtpModule,
6369        pairs: &[(&[f32], u32)],
6370        first_pos: usize,
6371    ) -> Result<(), &'static str> {
6372        for (j, (h, t)) in pairs.iter().enumerate() {
6373            self.mtp_warm(m, h, *t, first_pos + j);
6374        }
6375        Ok(())
6376    }
6377
6378    /// The MTP block alone — advance its KV with a (hidden, token) pair the
6379    /// verify just proved, without paying the head. What keeps the draft's
6380    /// attention context warm between speculative rounds.
6381    fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
6382        let e = self.embed_single(next_token);
6383        let mut cat = vec![0.0f32; 2 * self.hidden_size];
6384        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6385        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6386        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6387        let mut x = vec![0.0f32; self.hidden_size];
6388        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6389        inference::rms_norm_into(
6390            &x,
6391            &m.layer.input_norm,
6392            self.rms_eps,
6393            self.norm_style,
6394            &mut self.ws.n1,
6395        );
6396        let attn = match &m.layer.attn {
6397            AttnKind::Full {
6398                wq,
6399                wk,
6400                wv,
6401                wo,
6402                q_norm,
6403                k_norm,
6404                output_gate,
6405                softplus_gate,
6406                bias,
6407            } => {
6408                let mut cfg = self.attn_cfg(position);
6409                cfg.q_norm = q_norm.as_deref();
6410                cfg.k_norm = k_norm.as_deref();
6411                cfg.output_gate = *output_gate;
6412                cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
6413                cfg.bias = bias
6414                    .as_ref()
6415                    .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
6416                attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
6417            }
6418            _ => return,
6419        };
6420        let _ = attn;
6421    }
6422
6423    /// Speculative decode ON the wgpu whole-token graph: draft k with the
6424    /// MTP head, verify all of them plus the tip in ONE batched graph
6425    /// submit whose tail folds the head, commit the accepted prefix and
6426    /// roll the GDN state back to the last real position. Greedy only —
6427    /// output equals the plain graph's token for token, the way the DSV4
6428    /// verify equals the walk.
6429    #[cfg(feature = "gpu")]
6430    #[allow(clippy::too_many_arguments)]
6431    fn graph_spec_step(
6432        &mut self,
6433        m: &mut MtpModule,
6434        hidden: &[f32],
6435        t_next: u32,
6436        next_pos: usize,
6437        drafted: &mut usize,
6438        accepted: &mut usize,
6439        // The committed stream (prompt + generated so far, `t_next`
6440        // included): the sampler chain's penalties read it, and the
6441        // sampling arm extends it with the drafts position by position.
6442        all_ids: &mut Vec<u32>,
6443        // Tokens left before `max_tokens`. A round commits up to k
6444        // accepted drafts, and those positions are already in the cache,
6445        // so the depth is capped here — trimming the output afterwards
6446        // would leave cache rows the committed stream does not have.
6447        room: usize,
6448    ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
6449        // 3 is the measured optimum on Qwen3.6-27B / RTX 5090 (medians
6450        // of three, greedy): 51.1 tok/s against a plain 49.4, where k=2
6451        // gives 46.1, k=4 50.0, k=5 47.4, k=6 45.2. Acceptance is 89-91%
6452        // throughout — what turns the curve over is the verify, which
6453        // costs ~7.4 ms per extra position, and the draft ~3 ms a step.
6454        // 4 since the draft moved onto the graph (Qwen3.8-27B / 5090:
6455        // k=3 51.2, k=4 51.8 with the per-op draft; the graph draft
6456        // halves the draft cost, so the extra draft is cheaper still).
6457        // 5 with the int8 verify (the default: measured 76.5 against
6458        // k=4's 72-74 and k=6's 74 on the 5090), 4 with the f32 one.
6459        #[cfg(target_os = "macos")]
6460        let metal_native = crate::gpu::q1_force();
6461        #[cfg(not(target_os = "macos"))]
6462        let metal_native = false;
6463        #[cfg(feature = "gpu")]
6464        let k_default = if metal_native {
6465            // the Metal verify's GEMM tile is 8 rows wide and flat in b:
6466            // seven drafts + the tip fill it for free
6467            7
6468        } else if crate::gpu_wgpu::verify_i8_on() {
6469            5
6470        } else {
6471            4
6472        };
6473        #[cfg(not(feature = "gpu"))]
6474        let k_default = 4;
6475        let k_env: Option<usize> = std::env::var("CMF_GRAPH_SPEC_K")
6476            .ok()
6477            .and_then(|v| v.parse().ok())
6478            .filter(|&v| (1..=8).contains(&v));
6479        // Adaptive depth: start below the card's flat-verify optimum and
6480        // let the accepted fraction move it — predictable text climbs to
6481        // the old default within a few rounds, prose settles at 2-3 where
6482        // the shorter verify pays.
6483        let (k_start, k_max) = if metal_native { (7, 7) } else { (3, k_default.max(5)) };
6484        let k_full: usize = k_env.unwrap_or_else(|| self.spec_k_adapt.unwrap_or(k_start));
6485        let k_spec = k_full.min(room).max(1);
6486        // a tail round cut short by `room` says nothing about the text:
6487        // it must not move the adaptive depth the next request starts at
6488        let k_capped = k_spec < k_full;
6489        if next_pos == 0 {
6490            return None;
6491        }
6492        let t_round = std::time::Instant::now();
6493        // Submissions per phase — and they say where the round's money is.
6494        // Qwen3.6-27B on an RTX 5090, k=3:
6495        //
6496        //   draft   9.3 ms / 12 submissions   (four per MTP step)
6497        //   verify 52.8 ms /  1               (the batched graph)
6498        //   commit  5.4 ms /  6               (two per warm)
6499        //
6500        // The verify is already one submit. The draft's own work is 834 MB
6501        // a step — 0.8 ms at this card's measured 1056 GB/s — against 3.1
6502        // ms measured, so ~0.58 ms of every step is round trip, not
6503        // arithmetic, and the same holds for the warms. Eighteen round
6504        // trips a round at roughly half a millisecond each is ~11 ms of a
6505        // 68 ms round: fusing the MTP block into ONE submit the way the
6506        // trunk already is projects to ~64 tok/s against today's 50.9.
6507        // That is the largest measured item left on this path.
6508        let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
6509        let sub0 = subs();
6510        // Greedy without penalties verifies by argmax equality (bit-exact
6511        // against the plain path). Anything else is speculative SAMPLING:
6512        // each draft is a DRAW from the MTP head's post-chain distribution
6513        // q_j, kept for the accept test; the verify's rows give p_j.
6514        let cfg = self.sampler_config.clone();
6515        let penalized = !(cfg.repetition_penalty == 1.0
6516            && cfg.presence_penalty == 0.0
6517            && cfg.suppress_tokens.is_empty());
6518        // Three verify regimes: plain greedy (argmax of the raw rows),
6519        // greedy WITH penalties (argmax of the penalized rows — a single
6520        // pass each, no distributions), and sampling (draw / accept /
6521        // correct on post-chain distributions).
6522        let greedy_pen = cfg.temperature < 1e-6 && penalized;
6523        let sampling = cfg.temperature >= 1e-6;
6524        // Sampling with a top-k goes through the SPARSE chain: the dense
6525        // one builds nine 248k-float distributions a round (four drafts,
6526        // five verify rows) and measured 19-22 tok/s against a plain 40 —
6527        // the host, not the card. Sparse, the same nine cost tens of
6528        // microseconds each.
6529        let sparse = sampling && sampler::sparse_ok(&cfg);
6530        let base_len = all_ids.len();
6531        if sampling && !sparse && self.spec_q.len() < k_spec {
6532            self.spec_q.resize_with(k_spec, Vec::new);
6533        }
6534        if sparse && self.spec_qs.len() < k_spec {
6535            self.spec_qs.resize_with(k_spec, Vec::new);
6536        }
6537        // Draft the chain: first from the trunk's tip hidden, then the head
6538        // iterating on itself. Rows land in the MTP KV; the chain rows past
6539        // the first are speculation over speculative state and roll back
6540        // below, replaced by verified pairs.
6541        let mut drafts = Vec::with_capacity(k_spec);
6542        let mut hx = hidden.to_vec();
6543        // CMF_SPEC_DBG=1: draft 0 through BOTH MTP arms (graph and per-op)
6544        // from the same inputs — are the arms the difference, or the inputs?
6545        let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
6546        spec_stamp("pro");
6547        // Plain greedy on native Metal: the whole chain as one command
6548        // buffer (device argmax + embedding gather between the steps).
6549        // A decline before commit hands the round to the per-step loop
6550        // below; a failure after commit is terminal, like any graph
6551        // failure after admission.
6552        #[cfg(target_os = "macos")]
6553        if metal_native && !sampling && !greedy_pen && self.mtp_graph_mode != Some(false) {
6554            match self.mtp_draft_chain_metal(m, hidden, t_next, next_pos - 1, k_spec) {
6555                Ok(ids) => {
6556                    self.mtp_graph_mode = Some(true);
6557                    drafts = ids;
6558                }
6559                Err(true) => {
6560                    tracing::error!("mtp Metal draft chain failed after commit");
6561                    self.clear_sequence_state();
6562                    self.graph_failed
6563                        .store(true, std::sync::atomic::Ordering::Relaxed);
6564                    self.cancel
6565                        .store(true, std::sync::atomic::Ordering::Relaxed);
6566                    return None;
6567                }
6568                Err(false) => {}
6569            }
6570        }
6571        for j in drafts.len()..k_spec {
6572            let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
6573            let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
6574            if spec_dbg {
6575                let saved = self.mtp_graph_mode;
6576                self.mtp_graph_mode = Some(false);
6577                let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6578                self.mtp_graph_mode = saved;
6579                if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6580                    return None;
6581                }
6582                m.kv.truncate_last(1);
6583                dbg_ref = Some(r);
6584            }
6585            let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6586            if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6587                return None;
6588            }
6589            if let Some((lg_cpu, h_cpu)) = dbg_ref {
6590                let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
6591                let dl = lg
6592                    .iter()
6593                    .zip(&lg_cpu)
6594                    .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6595                let dh = hj
6596                    .iter()
6597                    .zip(&h_cpu)
6598                    .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6599                eprintln!(
6600                    "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 {}",
6601                    next_pos - 1 + j,
6602                    sampler::argmax(&lg_cpu),
6603                    sampler::argmax(&lg),
6604                    n(&h_cpu),
6605                    n(&hj),
6606                    m.kv.seq_len
6607                );
6608            }
6609            let dj = if sparse {
6610                let mut q = std::mem::take(&mut self.spec_qs[j]);
6611                let ok = sampler::sparse_distribution_into(
6612                    &lg,
6613                    &cfg,
6614                    all_ids,
6615                    &mut self.sampler_scratch,
6616                    self.pool.as_deref(),
6617                    &mut q,
6618                );
6619                let d = if ok {
6620                    sampler::draw_sparse(&q, &mut self.rng)
6621                } else {
6622                    // everything filtered: the dense chain's greedy fallback
6623                    let t = sampler::argmax(&lg);
6624                    q.clear();
6625                    q.push((t, 1.0));
6626                    t
6627                };
6628                self.spec_qs[j] = q;
6629                all_ids.push(d);
6630                d
6631            } else if sampling {
6632                let mut q = std::mem::take(&mut self.spec_q[j]);
6633                sampler::distribution_into(
6634                    &lg,
6635                    &cfg,
6636                    all_ids,
6637                    &mut self.sampler_scratch,
6638                    self.pool.as_deref(),
6639                    &mut q,
6640                );
6641                let d = sampler::draw(&q, &mut self.rng);
6642                self.spec_q[j] = q;
6643                all_ids.push(d); // the next draft's penalties see this one
6644                d
6645            } else if greedy_pen {
6646                let d = sampler::argmax_penalized(
6647                    &lg,
6648                    &cfg,
6649                    all_ids,
6650                    &mut self.sampler_scratch,
6651                    self.pool.as_deref(),
6652                );
6653                all_ids.push(d);
6654                d
6655            } else {
6656                sampler::argmax(&lg)
6657            };
6658            attention::recycle_buf(&mut lg);
6659            drafts.push(dj);
6660            hx = hj;
6661            spec_stamp("d.pick");
6662        }
6663        all_ids.truncate(base_len);
6664        *drafted += k_spec;
6665        let t_draft = t_round.elapsed();
6666        let sub_draft = subs();
6667        // Verify batch: [t_next, d1 .. d_{k-1}] at next_pos.. — every row's
6668        // logits come back from the graph's own head.
6669        let b = k_spec + 1;
6670        let mut hiddens = vec![0.0f32; b * self.hidden_size];
6671        for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
6672            let e = self.embed_single(t);
6673            hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
6674        }
6675        let positions: Vec<usize> = (next_pos..next_pos + b).collect();
6676        spec_stamp("v.emb");
6677        let (lm_gw, lm_rows) = {
6678            let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6679            (
6680                crate::gpu::GraphW {
6681                    idx: i,
6682                    kind,
6683                    row_scale: rs,
6684                    data: &[],
6685                    prism: crate::gpu::GraphPrismOp::None,
6686                    affine: false,
6687                },
6688                self.weights.lm_head.rows(),
6689            )
6690        };
6691        let mut logits = Vec::new();
6692        let final_norm = self.weights.final_norm.clone();
6693        // Plain greedy on Metal: the b argmaxes come from the device
6694        // (`argmax_rows` after the head) and the 7.9 MB logits plane is
6695        // never read back — the round's decision needs only the ids, and
6696        // the loop top takes the last verified id as `spec_forced`, which
6697        // is exactly what its argmax of the row would give. The full rows
6698        // stay for anything that reads them: sampling, penalties,
6699        // confidence, the verify oracle, the logit dump.
6700        // `CMF_METAL_DEV_ARGMAX=0` keeps the host path.
6701        #[cfg(target_os = "macos")]
6702        let greedy_dev = metal_native
6703            && !sampling
6704            && !greedy_pen
6705            && !self.confidence_on
6706            && self.final_softcap.is_none()
6707            // The host acceptance argmax scans the WHOLE head row
6708            // (`lm_rows`), the sampler's own row only `vocab_size`: they
6709            // coincide exactly when the head has no padding rows, and
6710            // only then is the device argmax (which scores `vocab_size`)
6711            // bit-identical to both.
6712            && self.vocab_size == lm_rows
6713            && std::env::var_os("CMF_METAL_VERIFY_CHECK").is_none()
6714            && std::env::var_os("CMF_LOGIT_DUMP").is_none()
6715            && std::env::var("CMF_METAL_DEV_ARGMAX").as_deref() != Ok("0");
6716        #[cfg(not(target_os = "macos"))]
6717        let greedy_dev = false;
6718        let mut dev_ids: Vec<u32> = Vec::new();
6719        #[cfg(target_os = "macos")]
6720        let verify_outcome = if metal_native {
6721            let lm = self.weights.lm_head.q1_parts()?;
6722            let n_score = self.vocab_size.min(lm_rows);
6723            self.try_batch_graph_metal(
6724                &mut hiddens,
6725                &positions,
6726                b,
6727                Some((lm, &final_norm, &mut logits)),
6728                if greedy_dev {
6729                    Some((n_score, &mut dev_ids))
6730                } else {
6731                    None
6732                },
6733            )
6734        } else {
6735            self.try_batch_graph_wgpu(
6736                &mut hiddens,
6737                &positions,
6738                b,
6739                Some(crate::gpu::SpecTail {
6740                    lm: lm_gw,
6741                    lm_rows,
6742                    final_norm: &final_norm,
6743                    logits_out: &mut logits,
6744                }),
6745            )
6746        };
6747        #[cfg(not(target_os = "macos"))]
6748        let verify_outcome = self.try_batch_graph_wgpu(
6749            &mut hiddens,
6750            &positions,
6751            b,
6752            Some(crate::gpu::SpecTail {
6753                lm: lm_gw,
6754                lm_rows,
6755                final_norm: &final_norm,
6756                logits_out: &mut logits,
6757            }),
6758        );
6759        match verify_outcome {
6760            crate::gpu::BatchGraphOutcome::Completed => {}
6761            crate::gpu::BatchGraphOutcome::Declined => {
6762                // The verifier refused before admission.  Its draft MTP
6763                // rows are still device-resident, so rewind the separate
6764                // mirror before the caller takes the exact one-token path.
6765                m.kv.truncate_last(k_spec);
6766                if !metal_native && !self.rewind_mtp_graph_mirror(next_pos) {
6767                    self.clear_sequence_state();
6768                    self.graph_failed
6769                        .store(true, std::sync::atomic::Ordering::Relaxed);
6770                    self.cancel
6771                        .store(true, std::sync::atomic::Ordering::Relaxed);
6772                    tracing::error!("MTP graph mirror rewind failed after verify decline");
6773                }
6774                return None;
6775            }
6776            crate::gpu::BatchGraphOutcome::Failed => {
6777                // A failed batch may have advanced trunk/GDN state.  Clear
6778                // both mirrors and preserve the terminal outcome rather than
6779                // falling through to stale CPU state.
6780                self.clear_sequence_state();
6781                self.graph_failed
6782                    .store(true, std::sync::atomic::Ordering::Relaxed);
6783                self.cancel
6784                    .store(true, std::sync::atomic::Ordering::Relaxed);
6785                tracing::error!("MTP verify batch graph failed after admission");
6786                return None;
6787            }
6788        }
6789        // `CMF_METAL_VERIFY_CHECK=1`: run the same b tokens through the
6790        // plain per-token path and compare each row's argmax + logits with
6791        // the verify's — the bring-up oracle for the batched graph. The
6792        // plain forwards mutate the CPU state; it is snapshotted and put
6793        // back, and the K/V mirrors re-pointed, before the round goes on.
6794        #[cfg(target_os = "macos")]
6795        if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
6796            let snap: Vec<Vec<f32>> = self
6797                .kv_cache
6798                .layers
6799                .iter()
6800                .map(|l| l.linear_state.clone())
6801                .collect();
6802            let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
6803            let toks: Vec<u32> = std::iter::once(t_next)
6804                .chain(drafts.iter().copied())
6805                .collect();
6806            let want_save = self.graph_want_logits;
6807            self.graph_want_logits = false;
6808            for (i, &t) in toks.iter().enumerate() {
6809                let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
6810                let _ = self.graph_logits.take();
6811                // CMF_SPEC_PLAIN_HIDDEN=1: the next round drafts from the
6812                // plain path's hidden instead of the verify's (an experiment
6813                // on the chain's sensitivity to the half-GEMM noise)
6814                if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
6815                    hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
6816                }
6817                let ref_lg = self.logits_from_hidden(&hi);
6818                let row = &logits[i * lm_rows..(i + 1) * lm_rows];
6819                let ra = sampler::argmax(&ref_lg);
6820                let va = sampler::argmax(row);
6821                let mut md = 0f32;
6822                let mut rms = 0f64;
6823                for j in 0..lm_rows.min(ref_lg.len()) {
6824                    let d = (ref_lg[j] - row[j]).abs();
6825                    md = md.max(d);
6826                    rms += (d as f64) * (d as f64);
6827                }
6828                let mut hd = 0f32;
6829                for j in 0..self.hidden_size {
6830                    hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
6831                }
6832                eprintln!(
6833                    "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
6834                    next_pos + i,
6835                    if ra == va { "OK" } else { "MISMATCH" },
6836                    (rms / lm_rows as f64).sqrt()
6837                );
6838            }
6839            self.graph_want_logits = want_save;
6840            // restore IN PLACE: the pending verify graph wraps these very
6841            // allocations (zero-copy) — replacing the Vec would strand it
6842            for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
6843                if l.linear_state.len() == st.len() {
6844                    l.linear_state.copy_from_slice(&st);
6845                } else {
6846                    l.linear_state = st;
6847                }
6848            }
6849            for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
6850                let extra = l.seq_len.saturating_sub(n0);
6851                if extra > 0 {
6852                    l.truncate_last(extra);
6853                    crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
6854                }
6855            }
6856        }
6857        let t_verify = t_round.elapsed();
6858        let sub_verify = subs();
6859        // Acceptance. Greedy: row i's argmax is the trunk's token after
6860        // input i. Sampling: accept draft i with min(1, p_i/q_i), and on
6861        // the first rejection draw the correction from max(0, p_i − q_i)
6862        // — that token is committed by the loop top as-is (spec_forced).
6863        let mut a = 0usize;
6864        let mut forced: Option<u32> = None;
6865        let ids: Vec<u32> = if sparse {
6866            let mut p = std::mem::take(&mut self.spec_ps);
6867            let mut res = std::mem::take(&mut self.spec_ress);
6868            while a < k_spec {
6869                let ok = sampler::sparse_distribution_into(
6870                    &logits[a * lm_rows..(a + 1) * lm_rows],
6871                    &cfg,
6872                    all_ids,
6873                    &mut self.sampler_scratch,
6874                    self.pool.as_deref(),
6875                    &mut p,
6876                );
6877                if !ok {
6878                    let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
6879                    p.clear();
6880                    p.push((t, 1.0));
6881                }
6882                match sampler::spec_accept_or_correct_sparse(
6883                    &p,
6884                    &self.spec_qs[a],
6885                    drafts[a],
6886                    &mut self.rng,
6887                    &mut res,
6888                ) {
6889                    None => {
6890                        all_ids.push(drafts[a]);
6891                        a += 1;
6892                    }
6893                    Some(c) => {
6894                        forced = Some(c);
6895                        break;
6896                    }
6897                }
6898            }
6899            all_ids.truncate(base_len);
6900            self.spec_ps = p;
6901            self.spec_ress = res;
6902            drafts.clone()
6903        } else if sampling {
6904            let mut p = std::mem::take(&mut self.spec_p);
6905            let mut res = std::mem::take(&mut self.spec_res);
6906            while a < k_spec {
6907                sampler::distribution_into(
6908                    &logits[a * lm_rows..(a + 1) * lm_rows],
6909                    &cfg,
6910                    all_ids,
6911                    &mut self.sampler_scratch,
6912                    self.pool.as_deref(),
6913                    &mut p,
6914                );
6915                match sampler::spec_accept_or_correct(
6916                    &p,
6917                    &self.spec_q[a],
6918                    drafts[a],
6919                    &mut self.rng,
6920                    &mut res,
6921                    self.pool.as_deref(),
6922                ) {
6923                    None => {
6924                        all_ids.push(drafts[a]);
6925                        a += 1;
6926                    }
6927                    Some(c) => {
6928                        forced = Some(c);
6929                        break;
6930                    }
6931                }
6932            }
6933            all_ids.truncate(base_len);
6934            self.spec_p = p;
6935            self.spec_res = res;
6936            // the accepted drafts ARE the verified tokens after inputs 0..a
6937            drafts.clone()
6938        } else if greedy_pen {
6939            // Row i's penalized argmax, penalties over the stream that
6940            // includes the accepted drafts before it — the plain loop's
6941            // exact arithmetic, one pass per row, no working copy.
6942            let mut ids: Vec<u32> = Vec::with_capacity(b);
6943            for i in 0..b {
6944                let t = sampler::argmax_penalized(
6945                    &logits[i * lm_rows..(i + 1) * lm_rows],
6946                    &cfg,
6947                    all_ids,
6948                    &mut self.sampler_scratch,
6949                    self.pool.as_deref(),
6950                );
6951                ids.push(t);
6952                if i < k_spec && t == drafts[i] {
6953                    all_ids.push(t);
6954                } else {
6955                    break;
6956                }
6957            }
6958            all_ids.truncate(base_len);
6959            while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
6960                a += 1;
6961            }
6962            // rows past the first mismatch were never scored; the loop
6963            // top re-samples the last verified row itself.
6964            ids
6965        } else if greedy_dev && dev_ids.len() == b {
6966            let ids = std::mem::take(&mut dev_ids);
6967            while a < k_spec && ids[a] == drafts[a] {
6968                a += 1;
6969            }
6970            ids
6971        } else {
6972            if logits.len() < b * lm_rows {
6973                // the device argmax was asked for and came back short:
6974                // no rows to fall back on — terminal like a failed batch
6975                self.clear_sequence_state();
6976                self.graph_failed
6977                    .store(true, std::sync::atomic::Ordering::Relaxed);
6978                self.cancel
6979                    .store(true, std::sync::atomic::Ordering::Relaxed);
6980                tracing::error!("Metal verify returned neither logits nor argmax ids");
6981                return None;
6982            }
6983            let ids: Vec<u32> = (0..b)
6984                .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
6985                .collect();
6986            while a < k_spec && ids[a] == drafts[a] {
6987                a += 1;
6988            }
6989            ids
6990        };
6991        spec_stamp("acc");
6992        if spec_dbg {
6993            eprintln!(
6994                "spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
6995                drafts, ids
6996            );
6997        }
6998        // CMF_METAL_VERIFY_CHECK=2: the commit oracle — plain-forward the
6999        // a+1 accepted tokens from a snapshot, then diff the replayed GDN
7000        // states and the appended K/V rows against that.
7001        #[cfg(target_os = "macos")]
7002        let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
7003            && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
7004        {
7005            let snap: Vec<Vec<f32>> = self
7006                .kv_cache
7007                .layers
7008                .iter()
7009                .map(|l| l.linear_state.clone())
7010                .collect();
7011            let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
7012            let toks: Vec<u32> = std::iter::once(t_next)
7013                .chain(drafts.iter().copied())
7014                .collect();
7015            let want_save = self.graph_want_logits;
7016            self.graph_want_logits = false;
7017            for (i, &t) in toks.iter().take(a + 1).enumerate() {
7018                let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
7019                let _ = self.graph_logits.take();
7020            }
7021            self.graph_want_logits = want_save;
7022            let plain_states: Vec<Vec<f32>> = self
7023                .kv_cache
7024                .layers
7025                .iter()
7026                .map(|l| l.linear_state.clone())
7027                .collect();
7028            let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7029            let mut rows = Vec::new();
7030            for (li, (l, n0)) in self
7031                .kv_cache
7032                .layers
7033                .iter_mut()
7034                .zip(attn_lens.iter())
7035                .enumerate()
7036            {
7037                let extra = l.seq_len.saturating_sub(*n0);
7038                if extra > 0 {
7039                    let mut kk = Vec::new();
7040                    let mut vv = Vec::new();
7041                    for g in 0..nkv {
7042                        kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7043                        vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7044                    }
7045                    rows.push((li, kk, vv));
7046                    l.truncate_last(extra);
7047                    crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
7048                }
7049            }
7050            for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7051                if l.linear_state.len() == st.len() {
7052                    l.linear_state.copy_from_slice(&st);
7053                } else {
7054                    l.linear_state = st;
7055                }
7056            }
7057            Some((plain_states, rows))
7058        } else {
7059            None
7060        };
7061        let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
7062        // Metal: the MTP cache cut and the round's warm-up SUBMIT come
7063        // BEFORE the trunk commit, so the warm-up's command buffer is
7064        // queued ahead of the GDN replay (second queue) and its wait
7065        // below no longer sits behind the replay — measured: the warm-up's
7066        // wait grew with the accepted count exactly like the replay does
7067        // (8 ms at a=1, 17 ms at a=3, 25 ms at a=5 for ~2 ms of its own
7068        // work). The replay now overlaps the warm-up's readback, the
7069        // round's return and the next draft chain.
7070        #[cfg(target_os = "macos")]
7071        let mut warm_pending: Option<MetalWarmPending> = None;
7072        #[cfg(target_os = "macos")]
7073        if metal_native {
7074            m.kv.truncate_last(k_spec.saturating_sub(1));
7075            if self.mtp_graph_mode == Some(true) {
7076                // the mirror rows below the cut are the CPU rows: re-point,
7077                // no re-upload
7078                crate::gpu_metal::kv_mirror_set_stored(
7079                    self.mtp_kv_id(),
7080                    Self::MTP_LAYER_BASE,
7081                    m.kv.seq_len,
7082                );
7083                if !warm_off && a > 0 {
7084                    let pairs: Vec<(&[f32], u32)> = (0..a)
7085                        .map(|j| {
7086                            (
7087                                &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
7088                                ids[j],
7089                            )
7090                        })
7091                        .collect();
7092                    warm_pending = self.mtp_warm_batch_submit(m, &pairs, next_pos);
7093                }
7094            }
7095            spec_stamp("c.wsub");
7096        }
7097        // a fully-accepted round needs no restore: every input was real.
7098        #[cfg(target_os = "macos")]
7099        if metal_native {
7100            // the Metal verify never wrote its states: the commit replays the
7101            // accepted prefix into the CPU owners and appends the K/V rows
7102            if !self.metal_verify_commit(a) {
7103                self.clear_sequence_state();
7104                self.graph_failed
7105                    .store(true, std::sync::atomic::Ordering::Relaxed);
7106                self.cancel
7107                    .store(true, std::sync::atomic::Ordering::Relaxed);
7108                tracing::error!("Metal verify state/KV handoff failed after admission");
7109                return None;
7110            }
7111            if let Some((plain_states, rows)) = commit_ref {
7112                crate::gpu_metal::queue_fence();
7113                // the commit's replay runs on the second queue: collect it
7114                // before the oracle reads the CPU owners it writes into
7115                let _ = crate::gpu_metal::wait_replay();
7116                let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7117                let mut worst_s = 0f32;
7118                let mut worst_li = 0usize;
7119                for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
7120                    if l.linear_state.len() != ps.len() || ps.is_empty() {
7121                        continue;
7122                    }
7123                    let d = l
7124                        .linear_state
7125                        .iter()
7126                        .zip(ps)
7127                        .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7128                    let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
7129                    let rel = d / n.max(1e-6);
7130                    if rel > worst_s {
7131                        worst_s = rel;
7132                        worst_li = li;
7133                    }
7134                }
7135                let mut worst_k = 0f32;
7136                for (li, kk, vv) in &rows {
7137                    let l = &self.kv_cache.layers[*li];
7138                    let n0 = l.seq_len - (kk.len() / (nkv * hd));
7139                    let mut ck = Vec::new();
7140                    let mut cv = Vec::new();
7141                    for g in 0..nkv {
7142                        ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7143                        cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7144                    }
7145                    if ck.len() == kk.len() {
7146                        let dk = ck
7147                            .iter()
7148                            .zip(kk)
7149                            .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7150                        let dv = cv
7151                            .iter()
7152                            .zip(vv)
7153                            .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7154                        worst_k = worst_k.max(dk).max(dv);
7155                    } else {
7156                        eprintln!(
7157                            "commit-check L{li}: kv row count mismatch {} vs {}",
7158                            ck.len(),
7159                            kk.len()
7160                        );
7161                    }
7162                }
7163                eprintln!(
7164                    "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}"
7165                );
7166            }
7167        }
7168        if !metal_native && a + 1 < b {
7169            let expected_gdn_layers = self.graph_gdn_layer_count();
7170            if expected_gdn_layers > 0
7171                && !crate::gpu::gdn_spec_restore(self.graph_kv_id, a, next_pos, expected_gdn_layers)
7172            {
7173                self.clear_sequence_state();
7174                self.graph_failed
7175                    .store(true, std::sync::atomic::Ordering::Relaxed);
7176                self.cancel
7177                    .store(true, std::sync::atomic::Ordering::Relaxed);
7178                tracing::error!("GDN speculative restore failed after verify");
7179                return None;
7180            }
7181        }
7182        if !metal_native && !self.rewind_trunk_graph_mirrors(next_pos + a + 1) {
7183            // The verify graph committed the full batch, but one of its
7184            // persistent Full-attention mirrors could not be re-pointed to
7185            // the accepted prefix.  Treat that as terminal state failure;
7186            // an exact CPU fallback would otherwise consume stale GDN/KV.
7187            self.clear_sequence_state();
7188            self.graph_failed
7189                .store(true, std::sync::atomic::Ordering::Relaxed);
7190            self.cancel
7191                .store(true, std::sync::atomic::Ordering::Relaxed);
7192            tracing::error!("trunk graph KV rewind failed after speculative verify");
7193            return None;
7194        }
7195        *accepted += a;
7196        // MTP cache: keep the first draft row (its inputs were real), drop
7197        // the chain's, then append the verified pairs the round produced.
7198        // Each of those is a whole MTP block on the per-op path and they
7199        // cost 5.8 ms of a 69 ms round at k=3 — a third of what the
7200        // round's own draft costs. PRICED, and they earn it: skipping
7201        // them (`CMF_SPEC_WARM=0`) drops acceptance from 89% to 81% at
7202        // k=3 and 85% to 74% at k=4, and the tok/s goes nowhere at k=3
7203        // (50.3 against 50.5) and backwards at k=4 (48.1 against 50.1).
7204        // The knob stays so the next person can re-price it after the
7205        // warms are batched instead of assuming either way.
7206        if !metal_native {
7207            // (Metal cut its MTP cache before the trunk commit, above)
7208            m.kv.truncate_last(k_spec.saturating_sub(1));
7209        }
7210        spec_stamp("c.trunc");
7211        if !metal_native
7212            && self.mtp_graph_mode == Some(true)
7213            && !self.rewind_mtp_graph_mirror(next_pos)
7214        {
7215            // The graph draft was admitted, so inability to move its cursor
7216            // back to the real anchor is a state failure, not a capability
7217            // refusal.  Do not warm or continue with a stale mirror.
7218            self.clear_sequence_state();
7219            self.graph_failed
7220                .store(true, std::sync::atomic::Ordering::Relaxed);
7221            self.cancel
7222                .store(true, std::sync::atomic::Ordering::Relaxed);
7223            tracing::error!("MTP graph mirror rewind failed after verify commit");
7224            return None;
7225        }
7226        if !warm_off && a > 0 {
7227            // Graph arm: all accepted pairs in ONE batched run over the
7228            // MTP block; the token graph one by one if the batch declines.
7229            let mut warmed = false;
7230            #[cfg(target_os = "macos")]
7231            if metal_native && self.mtp_graph_mode == Some(true) {
7232                // the batched warm-up was submitted before the trunk
7233                // commit: collect it here; one by one on the token graph
7234                // if it declined (or failed)
7235                warmed = match warm_pending.take() {
7236                    Some(p) => self.mtp_warm_batch_finish(m, p),
7237                    None => false,
7238                };
7239                if !warmed {
7240                    warmed = true;
7241                    for j in 0..a {
7242                        let row =
7243                            hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
7244                        if self
7245                            .mtp_step_metal(m, &row, ids[j], next_pos + j, false)
7246                            .is_none()
7247                        {
7248                            warmed = false;
7249                            break;
7250                        }
7251                    }
7252                }
7253            }
7254            if !warmed && self.mtp_graph_mode != Some(false) && !metal_native {
7255                let rows: Vec<Vec<f32>> = (0..a)
7256                    .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
7257                    .collect();
7258                let pairs: Vec<(&[f32], u32)> = rows
7259                    .iter()
7260                    .zip(ids.iter())
7261                    .map(|(r, &t)| (r.as_slice(), t))
7262                    .collect();
7263                match self.mtp_warm_prefill_pairs(m, &pairs, next_pos) {
7264                    Ok(()) => warmed = true,
7265                    Err(err) => {
7266                        // A warm-up failure after graph admission cannot
7267                        // fall back to `mtp_warm`: the detached CPU cache is
7268                        // not authoritative for the device mirror.  Mark it
7269                        // terminal so the generation caller clears state and
7270                        // returns instead of drafting from stale attention.
7271                        tracing::error!("{err}");
7272                        self.clear_sequence_state();
7273                        self.graph_failed
7274                            .store(true, std::sync::atomic::Ordering::Relaxed);
7275                        self.cancel
7276                            .store(true, std::sync::atomic::Ordering::Relaxed);
7277                        return None;
7278                    }
7279                }
7280            }
7281            if !warmed {
7282                for j in 0..a {
7283                    let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
7284                    let row = row.to_vec();
7285                    self.mtp_warm(m, &row, ids[j], next_pos + j);
7286                }
7287            }
7288        }
7289        // The sampler's contract: logits of the LAST verified position —
7290        // unless a rejected draft already drew the correction, in which
7291        // case the loop top commits that token and samples nothing.
7292        spec_stamp("c.warm");
7293        if let Some(c) = forced {
7294            self.spec_forced = Some(c);
7295            self.graph_logits = None;
7296        } else if greedy_dev && logits.is_empty() {
7297            // the row's argmax IS the token the loop top would pick from
7298            // it (plain greedy, no penalties): commit it as forced
7299            self.spec_forced = Some(ids[a]);
7300            self.graph_logits = None;
7301        } else {
7302            let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
7303            row.resize(self.vocab_size, 0.0);
7304            if let Some(c) = self.final_softcap {
7305                for l in row.iter_mut() {
7306                    *l = c * (*l / c).tanh();
7307                }
7308            }
7309            self.graph_logits = Some(row);
7310        }
7311        let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
7312        spec_stamp("c.row");
7313        // Three phases, not two. The round's wall clock was 4 ms longer
7314        // than draft+verify and the difference had nowhere to be seen:
7315        // the accepted prefix re-runs the MTP block once per token to
7316        // keep the draft head's attention cache warm, and the GDN state
7317        // rolls back on any rejection. Both live here, after the verify.
7318        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7319            let end = subs();
7320            eprintln!(
7321                "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
7322                 commit {:.1} ms/{} sub (accepted {a} of {k_spec}, full-head streak {})",
7323                t_draft.as_secs_f64() * 1e3,
7324                sub_draft - sub0,
7325                (t_verify - t_draft).as_secs_f64() * 1e3,
7326                sub_verify - sub_draft,
7327                (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
7328                end - sub_verify,
7329                self.draft_full_streak,
7330            );
7331        }
7332        // Native Metal's verify tile is flat in b (eight rows for the price
7333        // of one), so a shorter round only forfeits tokens — measured on
7334        // the M4: an essay round at k=2 still verified in 260 ms. The
7335        // adaptation is for cards whose verify grows with the rows.
7336        if k_env.is_none() && !metal_native && !k_capped {
7337            // Slow average and a wide band: a fast one oscillated 2↔3 on
7338            // an essay every other round (measured), which forfeits the
7339            // draft it just paid for.
7340            let f = a as f32 / k_spec.max(1) as f32;
7341            self.spec_acc_ewma += 0.2 * (f - self.spec_acc_ewma);
7342            let mut k_next = k_spec;
7343            if self.spec_acc_ewma >= 0.75 && k_spec < k_max {
7344                k_next = k_spec + 1;
7345            } else if self.spec_acc_ewma < 0.4 && k_spec > 2 {
7346                k_next = k_spec - 1;
7347            }
7348            if k_next != k_spec {
7349                self.spec_acc_ewma = 0.6;
7350                if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7351                    eprintln!("spec-k: {k_spec} → {k_next}");
7352                }
7353            }
7354            self.spec_k_adapt = Some(k_next);
7355        }
7356        spec_stamp("end");
7357        Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
7358    }
7359
7360    /// Micro-benchmark: two single-position forwards vs one fused pair
7361    /// from the current cache state (KV rewound after each probe).
7362    /// Returns (two_singles_ms, fused_pair_ms) per probe, or the (0, 0)
7363    /// sentinel when this model has no pair path to measure — the same
7364    /// answer the o1 arm gives, and the bench prints it the same way.
7365    /// (An architecture that loads its own layers leaves `weights.layers`
7366    /// empty; walking it here was an index panic, found by `bench` on
7367    /// deepseek_v4.)
7368    pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
7369        if !self.pair_supported() {
7370            return (0.0, 0.0);
7371        }
7372        // This is a host-side pair micro-benchmark. It truncates the host KV
7373        // after every probe, so letting the whole-token graph participate
7374        // would leave its device GDN/KV mirror ahead of the next probe and
7375        // poison the process-wide graph verdict before the real generation
7376        // benchmark starts. Keep the existing per-op/GPU arithmetic while
7377        // suppressing only the stateful token graph for this measurement.
7378        let graph_env = std::env::var_os("CMF_GPU_WGPU_GRAPH");
7379        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
7380        let emb1 = self.embed_single(1);
7381        let emb2 = self.embed_single(2);
7382        let pos = self.kv_cache.seq_len();
7383
7384        let t0 = std::time::Instant::now();
7385        for _ in 0..iters {
7386            let _ = self.forward_layers(&emb1, pos, None);
7387            let _ = self.forward_layers(&emb2, pos + 1, None);
7388            for l in &mut self.kv_cache.layers {
7389                l.truncate_last(2);
7390            }
7391        }
7392        let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7393
7394        let t1 = std::time::Instant::now();
7395        for _ in 0..iters {
7396            let _ = self.forward_pair(&emb1, &emb2, pos);
7397            for l in &mut self.kv_cache.layers {
7398                l.truncate_last(2);
7399            }
7400        }
7401        let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7402        match graph_env {
7403            Some(value) => unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", value) },
7404            None => unsafe { std::env::remove_var("CMF_GPU_WGPU_GRAPH") },
7405        }
7406        (singles_ms, pair_ms)
7407    }
7408
7409    /// Fused two-position forward: weight rows are streamed from memory
7410    /// once per layer for both positions. Full layers → fused GQA pair;
7411    /// linear layers → vmf_phase pair (lane 2 state is tentative in the
7412    /// per-layer scratch until the draft is accepted).
7413    /// Whether the fused two-position path covers every layer kind in
7414    /// this model. MLA and KDA run per position (their pair arms are
7415    /// unreachable); the seq prefill falls back to singles for them.
7416    fn pair_supported(&self) -> bool {
7417        // An EMPTY layer stack means the architecture loaded its own and
7418        // this path has nothing to walk. Checking that directly, rather
7419        // than naming each such architecture, is what makes the guard hold
7420        // for the next one: `any()` over no layers is false, so a
7421        // feature-by-feature test says "supported" for a model that has no
7422        // layers here at all.
7423        !self.weights.layers.is_empty()
7424            && self.g3n.is_none()
7425            && !self
7426                .weights
7427                .layers
7428                .iter()
7429                .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
7430    }
7431
7432    fn forward_pair(
7433        &mut self,
7434        emb1: &[f32],
7435        emb2: &[f32],
7436        position: usize,
7437    ) -> (Vec<f32>, Vec<f32>) {
7438        // A two-token prompt starts here, not in the layer walk: decide the
7439        // MiMo placement before the pair's per-op MoE uploads any expert.
7440        self.mimo_moe_prepare();
7441        let mut h1 = emb1.to_vec();
7442        let mut h2 = emb2.to_vec();
7443        let (_nkv, _hd, hs, _rd, eps) = (
7444            self.num_kv_heads,
7445            self.head_dim,
7446            self.hidden_size,
7447            self.rotary_dim,
7448            self.rms_eps,
7449        );
7450        let pool = self.pool.clone();
7451
7452        for li in 0..self.num_layers {
7453            let lw = &self.weights.layers[self.phys_layer(li)];
7454            // Norms into pipeline scratch (4 allocs/layer on the MTP
7455            // decode hot path before this).
7456            inference::rms_norm_into(
7457                &h1,
7458                &lw.input_norm,
7459                self.rms_eps,
7460                self.norm_style,
7461                &mut self.ws.n1,
7462            );
7463            inference::rms_norm_into(
7464                &h2,
7465                &lw.input_norm,
7466                self.rms_eps,
7467                self.norm_style,
7468                &mut self.ws.n2,
7469            );
7470
7471            let (a1, a2) = match &lw.attn {
7472                AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
7473                AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
7474                AttnKind::Bounded(w) => {
7475                    // Two sequential positions of the bounded operator
7476                    // (the ring's causal order is the pair's order).
7477                    let rope = self
7478                        .bounded_rope
7479                        .clone()
7480                        .expect("bounded layer without an installed rotation table");
7481                    let cfg = crate::bounded::BoundedAttnCfg {
7482                        num_heads: self.num_heads,
7483                        num_kv_heads: self.num_kv_heads,
7484                        head_dim: self.head_dim,
7485                        hidden_size: hs,
7486                        scale: self.attn_scale,
7487                        rope: &rope,
7488                        pool: pool.as_deref(),
7489                    };
7490                    let a1 = crate::bounded::bounded_attention(
7491                        &self.ws.n1,
7492                        w,
7493                        &mut self.kv_cache.layers[li],
7494                        &cfg,
7495                    );
7496                    let a2 = crate::bounded::bounded_attention(
7497                        &self.ws.n2,
7498                        w,
7499                        &mut self.kv_cache.layers[li],
7500                        &cfg,
7501                    );
7502                    (a1, a2)
7503                }
7504                AttnKind::Linear(w) => {
7505                    let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
7506                    let layer = &mut self.kv_cache.layers[li];
7507                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7508                    vmf_phase_pair(
7509                        &self.ws.n1,
7510                        &self.ws.n2,
7511                        w,
7512                        &cfg,
7513                        state,
7514                        scratch,
7515                        self.pool.as_deref(),
7516                    )
7517                }
7518                AttnKind::LinearGdn(w) => {
7519                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
7520                    let layer = &mut self.kv_cache.layers[li];
7521                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7522                    gdn_pair(
7523                        &self.ws.n1,
7524                        &self.ws.n2,
7525                        w,
7526                        &cfg,
7527                        state,
7528                        scratch,
7529                        self.pool.as_deref(),
7530                    )
7531                }
7532                AttnKind::ShortConv(w) => {
7533                    let cfg = self
7534                        .short_conv_cfg
7535                        .expect("short-conv layer without short_conv_cfg");
7536                    let layer = &mut self.kv_cache.layers[li];
7537                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7538                    short_conv_pair(
7539                        &self.ws.n1,
7540                        &self.ws.n2,
7541                        w,
7542                        &cfg,
7543                        state,
7544                        scratch,
7545                        self.pool.as_deref(),
7546                    )
7547                }
7548                AttnKind::Full {
7549                    wq,
7550                    wk,
7551                    wv,
7552                    wo,
7553                    q_norm,
7554                    k_norm,
7555                    output_gate,
7556                    softplus_gate,
7557                    bias,
7558                } => {
7559                    let inv_freq_l = self.layer_inv_freq(li);
7560                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
7561                    let cfg = QwenAttnCfg {
7562                        num_heads: self.layer_num_heads(li),
7563                        num_kv_heads: nkv_l,
7564                        head_dim: hd_l,
7565                        hidden_size: hs,
7566                        position,
7567                        inv_freq: &inv_freq_l,
7568                        rotary_dim: rd_l,
7569                        scale: self.attn_scale,
7570                        softcap: self.attn_softcap,
7571                        window: self.layer_window(li),
7572                        v_norm: self.attn_v_norm,
7573                        qk_norm_after_rope: self.qk_norm_after_rope,
7574                        q_norm: q_norm.as_deref(),
7575                        k_norm: k_norm.as_deref(),
7576                        output_gate: *output_gate,
7577                        softplus_gate: softplus_gate
7578                            .as_ref()
7579                            .map(|(gate, per_head)| (gate, *per_head)),
7580                        rope_scale: self.layer_rope_scale(li),
7581                        bias: bias
7582                            .as_ref()
7583                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7584                        rms_eps: eps,
7585                        norm_style: self.norm_style,
7586                        pool: pool.as_deref(),
7587                        v_head_dim: self.layer_v_dim(li),
7588                    };
7589                    attention::qwen_attention_pair(
7590                        &self.ws.n1,
7591                        &self.ws.n2,
7592                        wq,
7593                        wk,
7594                        wv,
7595                        wo,
7596                        &mut self.kv_cache.layers[li],
7597                        &cfg,
7598                    )
7599                }
7600            };
7601            let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
7602                Some(w) => (
7603                    inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
7604                    inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
7605                ),
7606                None => (a1, a2),
7607            };
7608            for i in 0..self.hidden_size {
7609                h1[i] += a1[i];
7610                h2[i] += a2[i];
7611            }
7612            let (mut a1, mut a2) = (a1, a2);
7613            attention::recycle_buf(&mut a1);
7614            attention::recycle_buf(&mut a2);
7615
7616            let lw = &self.weights.layers[self.phys_layer(li)];
7617            inference::rms_norm_into(
7618                &h1,
7619                &lw.post_norm,
7620                self.rms_eps,
7621                self.norm_style,
7622                &mut self.ws.p1,
7623            );
7624            inference::rms_norm_into(
7625                &h2,
7626                &lw.post_norm,
7627                self.rms_eps,
7628                self.norm_style,
7629                &mut self.ws.p2,
7630            );
7631            let (f1, f2) = match &lw.ffn {
7632                // Dual-branch layers need the raw residuals — run the
7633                // two positions through the same fn decode uses.
7634                FfnKind::DenseMoe(dm) => (
7635                    dense_moe_ffn(
7636                        dm,
7637                        &self.ws.p1,
7638                        &h1,
7639                        self.rms_eps,
7640                        self.norm_style,
7641                        self.pool.as_deref(),
7642                    ),
7643                    dense_moe_ffn(
7644                        dm,
7645                        &self.ws.p2,
7646                        &h2,
7647                        self.rms_eps,
7648                        self.norm_style,
7649                        self.pool.as_deref(),
7650                    ),
7651                ),
7652                FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => (
7653                    moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p1, self.pool.as_deref()),
7654                    moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p2, self.pool.as_deref()),
7655                ),
7656                _ => ffn_forward_pair(
7657                    &lw.ffn,
7658                    &self.ws.p1,
7659                    &self.ws.p2,
7660                    self.pool.as_deref(),
7661                    None,
7662                ),
7663            };
7664            let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
7665                Some(w) => (
7666                    inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
7667                    inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
7668                ),
7669                None => (f1, f2),
7670            };
7671            for i in 0..self.hidden_size {
7672                h1[i] += f1[i];
7673                h2[i] += f2[i];
7674            }
7675            let (mut f1, mut f2) = (f1, f2);
7676            attention::recycle_buf(&mut f1);
7677            attention::recycle_buf(&mut f2);
7678            if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
7679                for i in 0..self.hidden_size {
7680                    h1[i] *= sc;
7681                    h2[i] *= sc;
7682                }
7683            }
7684            // Looped Transformer: apply final norm at the end of each loop iteration.
7685            if self.is_loop_end(li) && li + 1 < self.num_layers {
7686                h1 = inference::rms_norm(
7687                    &h1,
7688                    &self.weights.final_norm,
7689                    self.rms_eps,
7690                    self.norm_style,
7691                );
7692                h2 = inference::rms_norm(
7693                    &h2,
7694                    &self.weights.final_norm,
7695                    self.rms_eps,
7696                    self.norm_style,
7697                );
7698            }
7699        }
7700        // Real O(1) prefill pairs may also carry tentative lane-2 recurrent
7701        // state. Commit it before publishing the transition epoch so the
7702        // next serial/device row cannot observe a new attention epoch with an
7703        // old GDN state. Speculative pairs run only when O(1) is inactive and
7704        // retain their existing caller-controlled commit/rollback semantics.
7705        if self.o1_active() {
7706            self.commit_linear_scratch();
7707        }
7708        self.o1_progress();
7709        (h1, h2)
7710    }
7711
7712    /// Commit lane-2 linear states after an accepted draft.
7713    fn commit_linear_scratch(&mut self) {
7714        for layer in &mut self.kv_cache.layers {
7715            if !layer.linear_scratch.is_empty() {
7716                std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
7717                layer.linear_scratch.clear();
7718            }
7719        }
7720    }
7721
7722    /// Forward a full id sequence from a fresh cache and return the
7723    /// logits after the last position (golden-parity harness, bench).
7724    pub fn forward_ids(
7725        &mut self,
7726        ids: &[u32],
7727        task_mask: Option<&TaskMask>,
7728    ) -> Result<Vec<f32>, String> {
7729        #[cfg(target_os = "macos")]
7730        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7731        if ids.is_empty() {
7732            return Err("empty id sequence".to_string());
7733        }
7734        self.clear_sequence_state();
7735        self.check_forward_graph("forward_ids setup", 0)?;
7736        if task_mask.is_none() {
7737            self.o1_begin();
7738        }
7739        let mut hidden = vec![0.0f32; self.hidden_size];
7740        let mut pos = 0usize;
7741        if let Some(b) = &mut self.dsv41 {
7742            let pool = self.pool.clone();
7743            let mut logits = Vec::new();
7744            crate::dsv41::forward_chunk(
7745                &b.0,
7746                &b.1,
7747                &b.2,
7748                &mut b.3,
7749                ids,
7750                0,
7751                pool.as_deref(),
7752                &mut logits,
7753            );
7754            if let Err(err) = self.o1_seal_checked() {
7755                self.clear_sequence_state();
7756                return Err(err);
7757            }
7758            return Ok(logits);
7759        }
7760        // Same routing predicate generation uses. Two reasons it must be
7761        // the same one: (1) a GDN hybrid's recurrent state is GPU-
7762        // resident, and a batched CPU prefill would build it on the host
7763        // only — decode then reads buffers the prefill never wrote;
7764        // (2) bench times THIS function and calls the result "prefill",
7765        // so a different path here reports a number production never
7766        // sees (W2 on 2×5090: 8.7 tok/s reported against 125 real).
7767        if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
7768            // prefill-GEMM in chunks; only the last position's hidden is
7769            // needed. (o1-compatible: the batch path attends per position
7770            // through qwen_attention, which carries the collection hook.)
7771            let chunk = self.prefill_chunk();
7772            let hs = self.hidden_size;
7773            while pos < ids.len() {
7774                let end = (pos + chunk).min(ids.len());
7775                let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
7776                self.check_forward_graph("forward_ids batched prefill", end - 1)?;
7777                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
7778                pos = end;
7779            }
7780        }
7781        // Same guards as generation's prefill — INCLUDING the graph one.
7782        // The CPU pair walk was intercepting positions that the resident
7783        // token graph would have run itself: on a GDN hybrid over wgpu
7784        // that is 89 ms of host forward against 7 ms of device submit,
7785        // and it made prefill look 12× slower than it is (W2 on an RTX
7786        // 5090, ctx 512: 11.2 tok/s with the walk, 136.6 without).
7787        // CMF_PAIR=0 opts out; a model whose layers live outside
7788        // `weights.layers` has no pair walk to take.
7789        if task_mask.is_none()
7790            && !self.graph_prefill_preferred()
7791            && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
7792            && self.pair_supported()
7793        {
7794            while pos + 1 < ids.len() {
7795                let e1 = self.embed_single(ids[pos]);
7796                let e2 = self.embed_single(ids[pos + 1]);
7797                let (_, h2) = self.forward_pair(&e1, &e2, pos);
7798                self.check_forward_graph("forward_ids pair", pos + 1)?;
7799                self.commit_linear_scratch();
7800                hidden = h2;
7801                pos += 2;
7802            }
7803        }
7804        // Resident Embryo graph: the prompt in chunks of one submit each
7805        // (the same device state and logits as the per-position walk).
7806        if task_mask.is_none() && pos == 0 && ids.len() > 1 {
7807            if let Some(lg) = self.embryo_prefill_chunked(ids, 0) {
7808                self.graph_logits = Some(lg);
7809                hidden = vec![0.0; self.hidden_size];
7810                pos = ids.len();
7811            }
7812        }
7813        while pos < ids.len() {
7814            hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
7815            self.check_forward_graph("forward_ids", pos)?;
7816            pos += 1;
7817        }
7818        if let Some(logits) = self.graph_logits.take() {
7819            // Resident stacks already applied final norm and their head in
7820            // the same submit; do not run a second norm/head over the zero
7821            // hidden sentinel returned by forward_layers_span.
7822            if let Err(err) = self.o1_seal_checked() {
7823                self.clear_sequence_state();
7824                return Err(err);
7825            }
7826            return Ok(logits);
7827        }
7828        // Harness contract: after forward_ids the cache is decode-ready —
7829        // under o1 that means sealed (bench measures the seal as part of
7830        // prefill, honestly).
7831        if let Err(err) = self.o1_seal_checked() {
7832            self.clear_sequence_state();
7833            return Err(err);
7834        }
7835        let normed = inference::rms_norm(
7836            &hidden,
7837            &self.weights.final_norm,
7838            self.rms_eps,
7839            self.norm_style,
7840        );
7841        Ok(self.lm_head_forward(&normed))
7842    }
7843
7844    /// Run the V4.1 stack one token at a time and retain logits for every
7845    /// position. This is a diagnostic surface for comparing a converted
7846    /// checkpoint with a tokenwise reference implementation.
7847    #[doc(hidden)]
7848    pub fn dsv41_serial_logits(&mut self, ids: &[u32]) -> Result<Vec<Vec<f32>>, String> {
7849        #[cfg(target_os = "macos")]
7850        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7851        if ids.is_empty() {
7852            return Err("empty id sequence".to_string());
7853        }
7854        self.clear_sequence_state();
7855        self.dsv41
7856            .as_ref()
7857            .ok_or_else(|| "dsv41 serial logits require a DeepSeek-V4.1 model".to_string())?;
7858        self.o1_begin();
7859        let rows = {
7860            let pool = self.pool.clone();
7861            let b = self
7862                .dsv41
7863                .as_mut()
7864                .expect("dsv41 checked above; state cannot change during forward");
7865            let mut rows = Vec::with_capacity(ids.len());
7866            for (position, &id) in ids.iter().enumerate() {
7867                let mut logits = Vec::new();
7868                crate::dsv41::forward_token(
7869                    &b.0,
7870                    &b.1,
7871                    &b.2,
7872                    &mut b.3,
7873                    id,
7874                    position,
7875                    pool.as_deref(),
7876                    &mut logits,
7877                );
7878                rows.push(logits);
7879            }
7880            rows
7881        };
7882        self.o1_seal();
7883        Ok(rows)
7884    }
7885
7886    /// Teacher-forced perplexity over a token sequence (phase-C gate:
7887    /// honest quant comparisons instead of prompt vibes).
7888    ///
7889    /// Attention is EXACT even on a model whose layers are flagged for
7890    /// the O(1) kernel — scoring the backbone is the default on purpose
7891    /// (it is the yardstick). `nll_ids_o1` scores the CONVERTED model.
7892    pub fn ppl_ids(&mut self, ids: &[u32]) -> Result<f64, String> {
7893        let (nll, cnt) = self.nll_ids_from(ids, 0)?;
7894        Ok((nll / cnt.max(1) as f64).exp())
7895    }
7896
7897    /// DTG-MA calibration pass (Patent 2): run `ids` through the model
7898    /// (CPU path, per position) and return each layer's per-neuron
7899    /// activation mass Σ|silu(gate)·up| — the statistic the task-guided
7900    /// FFN mask is derived from.
7901    pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
7902        self.clear_sequence_state();
7903        FFN_PROBE.with(|p| {
7904            *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
7905        });
7906        crate::gpu::cpu_scope(|| {
7907            for (pos, &id) in ids.iter().enumerate() {
7908                let emb = self.embed_single(id);
7909                let _ = self.forward_layers(&emb, pos, None);
7910            }
7911        });
7912        self.clear_sequence_state();
7913        FFN_PROBE
7914            .with(|p| p.borrow_mut().take())
7915            .unwrap_or_default()
7916    }
7917
7918    /// `probe_ffn_mass` over the BATCHED prefill: same accumulator, one
7919    /// sweep instead of one forward per token. What makes the statistic
7920    /// affordable on a 27B.
7921    pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Result<Vec<Vec<f64>>, String> {
7922        if let Err(err) = self.nll_begin() {
7923            // A recorder can be left by a caller that was interrupted before
7924            // this request entered its scoring block.  Consume it even when
7925            // the preflight failure prevents initialization of a new one.
7926            let _ = FFN_PROBE.with(|p| p.borrow_mut().take());
7927            self.nll_end();
7928            return Err(err);
7929        }
7930        FFN_PROBE.with(|p| {
7931            *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
7932        });
7933        let result: Result<(), String> = (|| {
7934            for chunk in ids.chunks(256) {
7935                if chunk.len() < 2 {
7936                    continue;
7937                }
7938                self.nll_ids_masked(chunk, 0, None)?;
7939            }
7940            Ok(())
7941        })();
7942        self.nll_end();
7943        let probe = FFN_PROBE
7944            .with(|p| p.borrow_mut().take())
7945            .unwrap_or_default();
7946        match result {
7947            Ok(()) => Ok(probe),
7948            Err(err) => {
7949                drop(probe);
7950                Err(err)
7951            }
7952        }
7953    }
7954
7955    /// Teacher-forced PPL with a task mask active (sparse execution) —
7956    /// the quality gate for a DTG-MA-masked skill. Sequential per
7957    /// position: the batched prefill path is dense-only.
7958    pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> Result<f64, String> {
7959        self.nll_begin()?;
7960        let result: Result<f64, String> = (|| {
7961            let mut nll = 0f64;
7962            let mut cnt = 0usize;
7963            let mut hidden = vec![0f32; self.hidden_size];
7964            for (pos, &id) in ids.iter().enumerate() {
7965                if pos > 0 {
7966                    inference::rms_norm_into(
7967                        &hidden,
7968                        &self.weights.final_norm,
7969                        self.rms_eps,
7970                        self.norm_style,
7971                        &mut self.ws.n1,
7972                    );
7973                    let mut logits = self.lm_head_forward(&self.ws.n1);
7974                    let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
7975                    let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
7976                    let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
7977                    nll -= p.max(1e-300).ln();
7978                    cnt += 1;
7979                    attention::recycle_buf(&mut logits);
7980                }
7981                let emb = self.embed_single(id);
7982                hidden = self.forward_layers(&emb, pos, Some(mask));
7983                self.nll_check_graph("masked serial forward", pos)?;
7984                // Consume a possible graph logits side channel before the
7985                // next row.  Masked scoring normally disables that route,
7986                // but stale channel state must never survive a request.
7987                let _ = self.graph_logits.take();
7988            }
7989            Ok((nll / cnt.max(1) as f64).exp())
7990        })();
7991        self.nll_end();
7992        result
7993    }
7994
7995    /// Teacher-forced NLL sum + scored-token count over positions
7996    /// `start..len-1`, attention EXACT. Positions below `start` still
7997    /// run — they are the context — they are just not scored, so this
7998    /// pairs with `nll_ids_o1(ids, start)` over the very same tokens.
7999    ///
8000    /// Returning (nll, cnt) rather than a ppl is what lets a windowed
8001    /// caller combine windows before the exp, so every scored token
8002    /// weighs the same regardless of how the windows are cut.
8003    /// `nll_ids_from` with a task mask held active at every position.
8004    ///
8005    /// The batched prefill path does not thread masks, so this walks the
8006    /// per-position forward — slower, but it scores the file exactly the
8007    /// way `run --task` will serve it, which is the point of the gate
8008    /// that calls it. With `None` it defers to the fast path.
8009    /// Masked scoring rides the SAME batched sweep as unmasked scoring —
8010    /// the masked-inference fast path: `prefill_batch_masked` lands the
8011    /// per-visit FFN rows on the activations inside the fused arms. The
8012    /// per-position loop below remains only as the no-batch fallback.
8013    pub fn nll_ids_masked(
8014        &mut self,
8015        ids: &[u32],
8016        start: usize,
8017        task_mask: Option<&TaskMask>,
8018    ) -> Result<(f64, usize), String> {
8019        let task_mask = self.drop_open_mask(task_mask);
8020        self.nll_ids_inner(ids, start, task_mask)
8021    }
8022
8023    pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> Result<(f64, usize), String> {
8024        self.nll_ids_inner(ids, start, None)
8025    }
8026
8027    fn nll_ids_inner(
8028        &mut self,
8029        ids: &[u32],
8030        start: usize,
8031        task_mask: Option<&TaskMask>,
8032    ) -> Result<(f64, usize), String> {
8033        self.nll_begin()?;
8034        let result: Result<(f64, usize), String> = (|| {
8035            let mut nll = 0f64;
8036            let mut cnt = 0usize;
8037            // An unmasked quality run with the resident wgpu graph must score
8038            // the same stateful path used by generation.  The layer-major
8039            // GEMM prefill below is a valid CPU/GEMM oracle, but it seeds
8040            // neither the graph's device GDN state nor its device KV mirrors;
8041            // using it here would silently score a different execution.  Keep
8042            // masked scoring on the exact per-position path as before, and
8043            // let the serial arm below drive the graph-aware scorer.
8044            // Only native Metal has a fused graph lm_head contract.  Vulkan
8045            // and other graph backends may expose hidden state without the
8046            // optional logits side channel; preserve their established CPU
8047            // norm/head fallback instead of turning that valid route into a
8048            // hard missing-logits error.
8049            let (graph_quality, fused_head_quality) = nll_graph_policy(
8050                task_mask.is_none(),
8051                self.graph_prefill_preferred(),
8052                crate::gpu::q1_force(),
8053            );
8054            self.graph_head_required = fused_head_quality;
8055            self.graph_want_logits = fused_head_quality;
8056            #[cfg(target_os = "macos")]
8057            if graph_quality && std::env::var("CMF_METAL_BATCH_NLL").as_deref() != Ok("0") {
8058                match self.nll_batch_metal(ids, start) {
8059                    MetalBatchNllOutcome::Completed(nll, count) => {
8060                        return Ok((nll, count));
8061                    }
8062                    MetalBatchNllOutcome::Declined => {}
8063                    MetalBatchNllOutcome::Failed(err) => return Err(err),
8064                }
8065            }
8066            if self.can_prefill_batched() && !graph_quality {
8067                // prefill-GEMM: layer-major position chunks, lm_head batched
8068                // (254MB lm_head read once per chunk, not per position).
8069                // The layer chunk is large (grouping positions by MoE experts
8070                // wins with size), lm_head in sub-blocks (logit buffer
8071                // 32×vocab ≈ 32MB instead of 128×).
8072                const CHUNK: usize = 128;
8073                const LM_SUB: usize = 32;
8074                let n = ids.len().saturating_sub(1);
8075                let hs = self.hidden_size;
8076                let rows = self.weights.lm_head.rows();
8077                let mut pos = 0usize;
8078                let state_trace = std::env::var("CMF_STATE_TRACE").is_ok();
8079                while pos < n {
8080                    let end = (pos + CHUNK).min(n);
8081                    let bsz = end - pos;
8082                    let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
8083                    self.nll_check_graph("batched prefill", pos)?;
8084                    if state_trace && end % 256 == 0 {
8085                        self.trace_recurrent_state(end);
8086                    }
8087                    let mut k0 = 0usize;
8088                    while k0 < bsz {
8089                        let k1 = (k0 + LM_SUB).min(bsz);
8090                        let sb = k1 - k0;
8091                        // Sub-block entirely below the scored range: the KV
8092                        // it just built is all this pass needed from it.
8093                        if pos + k1 <= start {
8094                            k0 = k1;
8095                            continue;
8096                        }
8097                        let mut normed = vec![0.0f32; sb * hs];
8098                        for k in 0..sb {
8099                            let r = inference::rms_norm(
8100                                &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
8101                                &self.weights.final_norm,
8102                                self.rms_eps,
8103                                self.norm_style,
8104                            );
8105                            normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
8106                        }
8107                        let mut logits = vec![0.0f32; sb * rows];
8108                        self.weights
8109                            .lm_head
8110                            .matmat(&normed, sb, &mut logits, self.pool.as_deref());
8111                        for k in 0..sb {
8112                            if pos + k0 + k < start {
8113                                continue;
8114                            }
8115                            self.nll_check_graph("batched score row", pos + k0 + k)?;
8116                            let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
8117                            if let Some(mu) = self.logit_multiplier {
8118                                for v in lg.iter_mut() {
8119                                    *v *= mu;
8120                                }
8121                            }
8122                            // Gemma-class final-logit soft-capping: the
8123                            // decode paths apply it; scoring must too, or
8124                            // the uncapped softmax misprices every token.
8125                            if let Some(c) = self.final_softcap {
8126                                for v in lg.iter_mut() {
8127                                    *v = c * (*v / c).tanh();
8128                                }
8129                            }
8130                            // Cortiq Embryo hierarchical head: same correction
8131                            // the decode path applies (lm_head_forward).
8132                            if let Some(cm) = self.head_clusters.clone() {
8133                                self.hierarchical_head_logprobs(
8134                                    &normed[k * hs..(k + 1) * hs],
8135                                    &cm,
8136                                    lg,
8137                                );
8138                            }
8139                            let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
8140                            let target = ids[pos + k0 + k + 1] as usize;
8141                            let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8142                            let lse: f64 = lg
8143                                .iter()
8144                                .map(|&v| ((v - max) as f64).exp())
8145                                .sum::<f64>()
8146                                .ln()
8147                                + max as f64;
8148                            nll += lse - lg[target] as f64;
8149                            cnt += 1;
8150                            if std::env::var("CMF_PPL_TRACE").is_ok() {
8151                                let top = lg
8152                                    .iter()
8153                                    .enumerate()
8154                                    .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8155                                    .map(|(i, _)| i)
8156                                    .unwrap_or(0);
8157                                eprintln!(
8158                                    "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
8159                                    pos + k0 + k,
8160                                    target,
8161                                    lse - lg[target] as f64,
8162                                    top,
8163                                    lg[target],
8164                                    lg[top]
8165                                );
8166                            }
8167                        }
8168                        k0 = k1;
8169                    }
8170                    pos = end;
8171                }
8172                return Ok((nll, cnt));
8173            }
8174            for pos in 0..ids.len().saturating_sub(1) {
8175                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
8176                self.nll_check_graph("serial forward", pos)?;
8177                // Architectures whose head lives inside their own stack return
8178                // the logits out of band and a zero hidden — DeepSeek-V4 folds
8179                // its hyper-connection copies between the last layer and the
8180                // norm, so it cannot hand back a vector this loop could use.
8181                // Scoring the zeros gave a perplexity of exactly the vocabulary
8182                // size, which is a uniform distribution reported as a
8183                // measurement. `generate` already reads this channel.
8184                let out_of_band = self.graph_logits.take();
8185                if self.graph_head_required && out_of_band.is_none() {
8186                    METAL_GRAPH_HEAD_MISS.fetch_add(
8187                        1,
8188                        std::sync::atomic::Ordering::Relaxed,
8189                    );
8190                    return Err(format!(
8191                        "fused Metal graph head did not complete at NLL position {pos}"
8192                    ));
8193                }
8194                if pos < start {
8195                    continue;
8196                }
8197                let logits = match out_of_band {
8198                    Some(lg) => lg,
8199                    None => {
8200                        let normed = inference::rms_norm(
8201                            &hidden,
8202                            &self.weights.final_norm,
8203                            self.rms_eps,
8204                            self.norm_style,
8205                        );
8206                        // lm_head_forward applies the final-logit softcap itself
8207                        // — capping again here double-squashed gemma-class
8208                        // logits (tanh∘tanh) and reported a flattered ppl.
8209                        self.lm_head_forward(&normed)
8210                    }
8211                };
8212                let target = ids[pos + 1] as usize;
8213                let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8214                let lse: f64 = logits
8215                    .iter()
8216                    .map(|&v| ((v - max) as f64).exp())
8217                    .sum::<f64>()
8218                    .ln()
8219                    + max as f64;
8220                let tok_nll = lse - logits[target] as f64;
8221                if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8222                    let top = logits
8223                        .iter()
8224                        .enumerate()
8225                        .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8226                        .map(|(i, _)| i)
8227                        .unwrap_or(0);
8228                    eprintln!(
8229                        "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8230                        logits[target], logits[top]
8231                    );
8232                }
8233                nll += tok_nll;
8234                cnt += 1;
8235            }
8236            Ok((nll, cnt))
8237        })();
8238        self.nll_end();
8239        result
8240    }
8241
8242    /// Score one post-layer hidden with the same final norm/head path used by
8243    /// decode. Keeping this in one helper is important for the production
8244    /// batch scorer: its rows stop before the final norm, just like the
8245    /// per-position O(1) path below.
8246    fn nll_from_hidden(&mut self, hidden: &[f32], target: u32, pos: usize) -> f64 {
8247        let normed = inference::rms_norm(
8248            hidden,
8249            &self.weights.final_norm,
8250            self.rms_eps,
8251            self.norm_style,
8252        );
8253        // lm_head_forward applies the final-logit softcap itself — capping
8254        // again here double-squashed gemma-class logits in earlier scorers.
8255        let mut logits = self.lm_head_forward(&normed);
8256        let target = target as usize;
8257        let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8258        let lse: f64 = logits
8259            .iter()
8260            .map(|&v| ((v - max) as f64).exp())
8261            .sum::<f64>()
8262            .ln()
8263            + max as f64;
8264        let tok_nll = lse - logits[target] as f64;
8265        if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8266            let top = logits
8267                .iter()
8268                .enumerate()
8269                .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8270                .map(|(i, _)| i)
8271                .unwrap_or(0);
8272            eprintln!(
8273                "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8274                logits[target], logits[top]
8275            );
8276        }
8277        attention::recycle_buf(&mut logits);
8278        tok_nll
8279    }
8280
8281    /// `CMF_STATE_TRACE`: per-layer magnitude of the recurrent record at
8282    /// position `pos` — the whole `linear_state` (vmf: S then the conv
8283    /// ring; GDN: conv ring then S), its recurrent S part alone, the
8284    /// bounded ring, and the last per-position KV row count.  The tool
8285    /// that separated the ~4k perplexity cliff of the 500-step exports
8286    /// (a state that keeps climbing past the trained window) from a
8287    /// runtime boundary; one line per layer, `STATE pos=… layer=…`.
8288    fn trace_recurrent_state(&self, pos: usize) {
8289        let stats = |v: &[f32]| -> (f64, f64) {
8290            if v.is_empty() {
8291                return (0.0, 0.0);
8292            }
8293            let (mut ss, mut mx) = (0f64, 0f64);
8294            for &x in v {
8295                ss += (x as f64) * (x as f64);
8296                mx = mx.max((x as f64).abs());
8297            }
8298            ((ss / v.len() as f64).sqrt(), mx)
8299        };
8300        for (li, l) in self.kv_cache.layers.iter().enumerate() {
8301            let lw = &self.weights.layers[self.phys_layer(li)];
8302            let (kind, s_len) = match &lw.attn {
8303                AttnKind::Linear(_) => (
8304                    "vmf",
8305                    self.vmf_cfg.map(|c| c.state_len()).unwrap_or(0),
8306                ),
8307                AttnKind::LinearGdn(_) => (
8308                    "gdn",
8309                    self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0),
8310                ),
8311                AttnKind::Bounded(_) => ("bounded", 0),
8312                AttnKind::Full { .. } => ("full", 0),
8313                _ => ("other", 0),
8314            };
8315            let (rms, max) = stats(&l.linear_state);
8316            let s_part = if kind == "vmf" {
8317                &l.linear_state[..s_len.min(l.linear_state.len())]
8318            } else if kind == "gdn" {
8319                let ring = l.linear_state.len().saturating_sub(s_len).min(l.linear_state.len());
8320                &l.linear_state[ring..]
8321            } else {
8322                &l.linear_state[..0]
8323            };
8324            let (s_rms, s_max) = stats(s_part);
8325            let (ring_rms, ring_len) = match &l.bounded {
8326                Some(b) => (stats(&b.ring_k).0, b.ring_k.len()),
8327                None => (0.0, 0),
8328            };
8329            eprintln!(
8330                "STATE pos={pos} layer={li} kind={kind} state_len={} rms={rms:.5} max={max:.4} \
8331                 S_rms={s_rms:.5} S_max={s_max:.4} ring_k_rms={ring_rms:.5} ring_len={ring_len} kv_rows={}",
8332                l.linear_state.len(),
8333                l.seq_len
8334            );
8335        }
8336    }
8337
8338    /// Teacher-forced NLL of the CONVERTED model: the O(1) Nyström path
8339    /// is ACTIVE over the scored positions. Returns `Ok((nll sum, scored
8340    /// count))` over `prefill..len-1` and surfaces a post-mutation batch
8341    /// failure instead of returning a partial score.
8342    ///
8343    /// Runtime discipline, deliberately NOT the matrix probe's: the
8344    /// requested prefix plus any required deferred lead-in run the exact
8345    /// prompt pass — that pass is what freezes the landmarks and M — and
8346    /// every post-seal scored position goes through `NystromState::step()`,
8347    /// the same code decode runs.
8348    /// So the landmarks are PREFILL-frozen (what ships), not
8349    /// full-sequence oracles (what the published probe measured). When the
8350    /// requested prefix is shorter than the bounded transition, rows in the
8351    /// exact lead-in are still scored so the shifted target range is stable.
8352    ///
8353    /// Pair with `nll_ids_from(ids, prefill)` for the exact baseline
8354    /// over the identical token set — that ratio is the honest one.
8355    pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> Result<(f64, usize), String> {
8356        // This scorer consumes host hiddens, so never request the optional
8357        // token-graph lm_head side channel. `nll_begin` also consumes a
8358        // prior graph failure and clears only the cancel bit that failure
8359        // raised, leaving a caller-owned cancellation observable.
8360        self.nll_begin()?;
8361        let requested_prefix = (prefill > 0).then_some(prefill);
8362        self.o1_begin_with_prefix(requested_prefix);
8363        let n = ids.len().saturating_sub(1);
8364        let requested_start = prefill.min(n);
8365        // The exact prefix must reach the deferred boundary before a
8366        // collecting layer can convert. Rows between the requested start and
8367        // that boundary remain part of the public NLL range and are scored
8368        // from the same hidden pass below.
8369        let exact_end = if self.o1_active() {
8370            match requested_prefix {
8371                Some(requested) => self.o1_effective_boundary(requested),
8372                None => self
8373                    .o1_cfg
8374                    .as_ref()
8375                    .and_then(|c| crate::nystrom::o1_deferred_boundary(c.w, c.sink)),
8376            }
8377            .unwrap_or(requested_start)
8378            .min(n)
8379        } else {
8380            requested_start
8381        };
8382        let mut nll = 0f64;
8383        let mut cnt = 0usize;
8384
8385        // Exact prompt pass over ids[..exact_end]: the seal consumes its
8386        // q/k/v. Rows at or after requested_start are scored here when the
8387        // bounded lead-in is longer than the caller's requested prefix.
8388        let mut pos = 0usize;
8389        if self.can_prefill_batched() {
8390            const CHUNK: usize = 128;
8391            while pos < exact_end {
8392                let end = (pos + CHUNK).min(exact_end);
8393                let hiddens = self.prefill_batch(&ids[pos..end], pos);
8394                if self
8395                    .graph_failed
8396                    .swap(false, std::sync::atomic::Ordering::Relaxed)
8397                {
8398                    self.cancel
8399                        .store(false, std::sync::atomic::Ordering::Relaxed);
8400                    self.nll_end();
8401                    return Err("GPU graph failed during O(1) NLL prefix".into());
8402                }
8403                for row in 0..end - pos {
8404                    let score_pos = pos + row;
8405                    if score_pos >= requested_start && score_pos < n {
8406                        nll += self.nll_from_hidden(
8407                            &hiddens[row * self.hidden_size..(row + 1) * self.hidden_size],
8408                            ids[score_pos + 1],
8409                            score_pos,
8410                        );
8411                        cnt += 1;
8412                    }
8413                }
8414                pos = end;
8415            }
8416        } else {
8417            while pos < exact_end {
8418                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8419                if self
8420                    .graph_failed
8421                    .swap(false, std::sync::atomic::Ordering::Relaxed)
8422                {
8423                    self.cancel
8424                        .store(false, std::sync::atomic::Ordering::Relaxed);
8425                    self.nll_end();
8426                    return Err("GPU graph failed during O(1) NLL prefix".into());
8427                }
8428                if pos >= requested_start && pos < n {
8429                    nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8430                    cnt += 1;
8431                }
8432                pos += 1;
8433            }
8434        }
8435        self.o1_seal_checked().map_err(|err| {
8436            self.nll_end();
8437            err
8438        })?;
8439
8440        // Reuse the production whole-token batch graph for the post-seal
8441        // suffix when the caller explicitly enabled both routes. This is a
8442        // teacher-forced scorer, so every row is ids[pos] and its target is
8443        // ids[pos + 1]; no speculative tail or rollback state is involved.
8444        // A first Declined is safe to handle with the established serial O(1)
8445        // path. Once a chunk completes, however, the device recurrent state
8446        // owns the sequence and a later decline must be terminal rather than
8447        // falling back to stale CPU state.
8448        let batch_k = std::env::var("CMF_BATCH_K")
8449            .ok()
8450            .and_then(|v| v.parse::<usize>().ok())
8451            .unwrap_or(0);
8452        let batch_admitted = batch_k > 0
8453            && self.can_prefill_batched()
8454            && self.o1_active()
8455            && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
8456            && (0..self.num_layers).all(|li| {
8457                let cache = &self.kv_cache.layers[self.phys_layer(li)];
8458                cache.o1.is_none() || cache.o1_views().is_some()
8459            });
8460        if std::env::var("CMF_GRAPH_PROF").is_ok() {
8461            eprintln!(
8462                "nll-batch: phase=post-seal admission={} requested_k={} scored_rows={}",
8463                batch_admitted,
8464                batch_k,
8465                n.saturating_sub(exact_end),
8466            );
8467        }
8468        let mut batch_completed = false;
8469        if batch_admitted && exact_end < n {
8470            let hs = self.hidden_size;
8471            let mut batch_pos = exact_end;
8472            while batch_pos < n {
8473                let end = (batch_pos + batch_k).min(n);
8474                let bk = end - batch_pos;
8475                let mut hiddens = vec![0.0f32; bk * hs];
8476                for (row, &id) in ids[batch_pos..end].iter().enumerate() {
8477                    hiddens[row * hs..(row + 1) * hs].copy_from_slice(&self.embed_single(id));
8478                }
8479                let positions: Vec<usize> = (batch_pos..end).collect();
8480                let t_batch = std::time::Instant::now();
8481                let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
8482                if std::env::var("CMF_GRAPH_PROF").is_ok() {
8483                    let ms = t_batch.elapsed().as_secs_f64() * 1000.0;
8484                    eprintln!(
8485                        "nll-batch: phase=post-seal mode=o1 k={bk} pos={}..{} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
8486                        batch_pos,
8487                        end.saturating_sub(1),
8488                        bk as f64 / (ms / 1000.0),
8489                    );
8490                }
8491                if let Err(err) = self.nll_check_graph("batch graph", batch_pos) {
8492                    self.nll_end();
8493                    return Err(err);
8494                }
8495                match outcome {
8496                    crate::gpu::BatchGraphOutcome::Completed => {
8497                        batch_completed = true;
8498                        for row in 0..bk {
8499                            nll += self.nll_from_hidden(
8500                                &hiddens[row * hs..(row + 1) * hs],
8501                                ids[batch_pos + row + 1],
8502                                batch_pos + row,
8503                            );
8504                            cnt += 1;
8505                        }
8506                        batch_pos = end;
8507                    }
8508                    crate::gpu::BatchGraphOutcome::Declined => {
8509                        if batch_completed {
8510                            self.nll_end();
8511                            return Err(format!(
8512                                "O(1) NLL batch declined after completed chunk at position {batch_pos}"
8513                            ));
8514                        }
8515                        break;
8516                    }
8517                    crate::gpu::BatchGraphOutcome::Failed => {
8518                        self.nll_end();
8519                        return Err(format!(
8520                            "O(1) NLL batch graph failed after admission at position {batch_pos}"
8521                        ));
8522                    }
8523                }
8524            }
8525            if batch_completed && cnt == n.saturating_sub(requested_start) {
8526                self.nll_end();
8527                return Ok((nll, cnt));
8528            }
8529        }
8530
8531        // Serial O(1) fallback/reference. It is intentionally retained when
8532        // batch admission declines before mutation; callers must label this
8533        // CMF_BATCH_K=0/per-position path separately from the production
8534        // whole-token batch route.
8535        for pos in exact_end..n {
8536            let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8537            if self
8538                .graph_failed
8539                .swap(false, std::sync::atomic::Ordering::Relaxed)
8540            {
8541                self.cancel
8542                    .store(false, std::sync::atomic::Ordering::Relaxed);
8543                self.nll_end();
8544                return Err(format!(
8545                    "GPU graph failed during O(1) NLL serial scoring at position {pos}"
8546                ));
8547            }
8548            nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8549            cnt += 1;
8550        }
8551        self.nll_end();
8552        Ok((nll, cnt))
8553    }
8554
8555    /// Teacher-forced calibration data (B1): for each position, whether the
8556    /// argmax equals the actual next token, and the top-1 softmax prob
8557    /// (top-1 probability) under EACH temperature in `temps` — all from ONE forward
8558    /// pass (argmax/correctness are temperature-invariant; only p_max
8559    /// reshapes). Feeds `cortiq calibrate` (reliability/ECE + temperature
8560    /// fit): is the model's confidence a true property, or does it need a
8561    /// measured scaling?
8562    pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
8563        self.clear_sequence_state();
8564        let n = ids.len().saturating_sub(1);
8565        let mut correct = Vec::with_capacity(n);
8566        let mut pmax = Vec::with_capacity(n);
8567        for pos in 0..n {
8568            let emb = self.embed_single(ids[pos]);
8569            let hidden = self.forward_layers(&emb, pos, None);
8570            let logits = if let Some(logits) = self.graph_logits.take() {
8571                logits
8572            } else {
8573                let normed = inference::rms_norm(
8574                    &hidden,
8575                    &self.weights.final_norm,
8576                    self.rms_eps,
8577                    self.norm_style,
8578                );
8579                // lm_head_forward applies the final-logit softcap itself —
8580                // capping again here double-squashed gemma-class logits
8581                // (tanh∘tanh) and reported a flattered ppl.
8582                self.lm_head_forward(&normed)
8583            };
8584            let target = ids[pos + 1] as usize;
8585            let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
8586            for (i, &v) in logits.iter().enumerate() {
8587                if v > mval {
8588                    mval = v;
8589                    amax = i;
8590                }
8591            }
8592            correct.push(amax == target);
8593            let row: Vec<f32> = temps
8594                .iter()
8595                .map(|&t| {
8596                    let tt = t.max(1e-3);
8597                    let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
8598                    1.0 / s.max(1e-12) // numerator at the max is exp(0)=1
8599                })
8600                .collect();
8601            pmax.push(row);
8602        }
8603        self.clear_sequence_state();
8604        (correct, pmax)
8605    }
8606
8607    /// Teacher-forced PPL with the dynamic router driving per-window
8608    /// skill switches (VMF experiment №2 measurement). Sequential (φ
8609    /// must update per token), returns (ppl, switch_count). The router
8610    /// must be enabled (`enable_dynamic_routing`); else this equals
8611    /// plain `ppl_ids`. The active skill when scoring token t shapes the
8612    /// logits for t+1 — on-policy over the held-out text itself.
8613    pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> Result<(f64, usize), String> {
8614        if self.dyn_router.is_none() {
8615            return Ok((self.ppl_ids(ids)?, 0));
8616        }
8617        self.nll_begin()?;
8618        let saved_active = self.dyn_active;
8619        let mut router = self
8620            .dyn_router
8621            .take()
8622            .ok_or_else(|| "dynamic router disappeared before PPL scoring".to_string())?;
8623        router.reset();
8624        self.dyn_phi_seen = 0;
8625        let _ = self.set_active_skill(None);
8626
8627        let result: Result<(f64, usize), String> = (|| {
8628            let mut nll = 0f64;
8629            let mut cnt = 0usize;
8630            for pos in 0..ids.len().saturating_sub(1) {
8631                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8632                self.nll_check_graph("dynamic serial forward", pos)?;
8633                let out_of_band = self.graph_logits.take();
8634                let mut logits = match out_of_band {
8635                    Some(lg) => lg,
8636                    None => {
8637                        let normed = inference::rms_norm(
8638                            &hidden,
8639                            &self.weights.final_norm,
8640                            self.rms_eps,
8641                            self.norm_style,
8642                        );
8643                        // lm_head_forward applies the final-logit softcap itself —
8644                        // capping again here double-squashed gemma-class logits
8645                        // and reported a flattered ppl.
8646                        self.lm_head_forward(&normed)
8647                    }
8648                };
8649                let target = ids[pos + 1] as usize;
8650                let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8651                let lse: f64 = logits
8652                    .iter()
8653                    .map(|&v| ((v - max) as f64).exp())
8654                    .sum::<f64>()
8655                    .ln()
8656                    + max as f64;
8657                let tok_nll = lse - logits[target] as f64;
8658                if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8659                    let top = logits
8660                        .iter()
8661                        .enumerate()
8662                        .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8663                        .map(|(i, _)| i)
8664                        .unwrap_or(0);
8665                    eprintln!(
8666                        "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8667                        logits[target], logits[top]
8668                    );
8669                }
8670                nll += tok_nll;
8671                cnt += 1;
8672                attention::recycle_buf(&mut logits);
8673                // Route on the evolving phi (drives the NEXT token's skill).
8674                let phi = self.dyn_phi_ema.clone();
8675                if let Some(new_active) = router.step(&phi, pos) {
8676                    let _ = self.set_active_skill(new_active);
8677                }
8678            }
8679            Ok(((nll / cnt.max(1) as f64).exp(), router.switches.len()))
8680        })();
8681
8682        // Restore the detached router and the active overlay on both success
8683        // and failure. The scoring state is cleared independently below.
8684        let _ = self.set_active_skill(saved_active);
8685        self.dyn_router = Some(router);
8686        self.nll_end();
8687        result
8688    }
8689
8690    /// Routing probe φ (spec §9): mean-pooled hidden after `layer`.
8691    pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
8692        self.clear_sequence_state();
8693        let mut acc = vec![0f32; self.hidden_size];
8694        for (pos, &id) in ids.iter().enumerate() {
8695            let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8696            for (a, v) in acc.iter_mut().zip(&h) {
8697                *a += v;
8698            }
8699        }
8700        let n = ids.len().max(1) as f32;
8701        for a in acc.iter_mut() {
8702            *a /= n;
8703        }
8704        self.clear_sequence_state();
8705        acc
8706    }
8707
8708    /// Router-v2 φ probe (spec §9.4, `phi.pool = "span_mean"`): the hidden
8709    /// AFTER `layer` — the same per-position walk and the same quantity as
8710    /// [`Self::probe_phi`] — averaged over the positions in `span` only
8711    /// (the user text between the template's prefix and suffix ids), NOT
8712    /// unit-normalized (the decision normalizes). The walk stops at
8713    /// `span.end`: causality makes the later positions irrelevant, so the
8714    /// result is bit-identical to probing `ids[..span.end]`. An empty span
8715    /// gives the zero vector, which the decision treats as degenerate.
8716    ///
8717    /// Every sequence state is reset before and after — the host KV/ring/
8718    /// recurrent state, the reuse keys (`kv_history`, `kv_prefix`) and the
8719    /// device graph's sequence — so run it on a pipeline that does not
8720    /// also serve a conversation (its prefix reuse would be lost).
8721    pub fn probe_phi_span(
8722        &mut self,
8723        ids: &[u32],
8724        layer: usize,
8725        span: std::ops::Range<usize>,
8726    ) -> Vec<f32> {
8727        #[cfg(target_os = "macos")]
8728        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8729        let end = span.end.min(ids.len());
8730        let start = span.start.min(end);
8731        let reset = |p: &mut Self| p.clear_sequence_state();
8732        reset(self);
8733        let mut acc = vec![0f32; self.hidden_size];
8734        for (pos, &id) in ids[..end].iter().enumerate() {
8735            let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8736            if pos >= start {
8737                for (a, v) in acc.iter_mut().zip(&h) {
8738                    *a += v;
8739                }
8740            }
8741        }
8742        let n = end - start;
8743        if n > 0 {
8744            let n = n as f32;
8745            for a in acc.iter_mut() {
8746                *a /= n;
8747            }
8748        }
8749        reset(self);
8750        acc
8751    }
8752
8753    /// One decode step of the current sequence: forward `token` at
8754    /// `position` (the cache holds positions `[0, position)`, e.g. after
8755    /// [`Self::forward_ids`]) and return the next-token logits — the same
8756    /// forward and head the generation loop runs (resident-graph logits
8757    /// when the graph ran, final norm + lm_head otherwise). The logit-dump
8758    /// tools drive greedy decoding with it so every position is observable.
8759    pub fn decode_step_logits(&mut self, token: u32, position: usize) -> Vec<f32> {
8760        #[cfg(target_os = "macos")]
8761        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8762        self.graph_logits = None;
8763        let hidden = self.forward_layers(&self.embed_single(token), position, None);
8764        if let Some(logits) = self.graph_logits.take() {
8765            return logits;
8766        }
8767        inference::rms_norm_into(
8768            &hidden,
8769            &self.weights.final_norm,
8770            self.rms_eps,
8771            self.norm_style,
8772            &mut self.ws.n1,
8773        );
8774        self.lm_head_forward(&self.ws.n1)
8775    }
8776
8777    /// Layer-major batched prefill (prefill-GEMM): full-attention —
8778    /// per-position with the existing operators (KV grows naturally,
8779    /// causality preserved), GDN projections / FFN / MoE — batched
8780    /// (a weight row is read from DRAM once per chunk, not per
8781    /// position). Returns the hidden of all positions [b × hidden].
8782    fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
8783        self.prefill_batch_masked(ids, start_pos, None)
8784    }
8785
8786    /// `prefill_batch` with a task mask honored on the dense-FFN panels
8787    /// (the masked-inference fast path: full fused compute, mask lands on
8788    /// the activations). The whole-chunk GPU graph is skipped for masked
8789    /// layers by the callers' arms; the per-GEMM device paths stay in
8790    /// play because the zeroing happens on the host between them.
8791    fn prefill_batch_masked(
8792        &mut self,
8793        ids: &[u32],
8794        start_pos: usize,
8795        task_mask: Option<&TaskMask>,
8796    ) -> Vec<f32> {
8797        self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
8798    }
8799
8800    /// One prompt chunk through the whole stack, post-stack rows out (no
8801    /// final norm) — the ingest generation uses, shared by scoring and
8802    /// `forward_ids` so they measure the same execution: the batched wgpu
8803    /// graph's device prefix plus the host's batched walk for the rest when
8804    /// `batch_prefix_prefill` holds and the graph admits the chunk, else
8805    /// the host's chunked prefill. Err only when a graph that had mutated
8806    /// device state failed.
8807    fn prefill_rows(
8808        &mut self,
8809        ids: &[u32],
8810        pos: usize,
8811        task_mask: Option<&TaskMask>,
8812    ) -> Result<Vec<f32>, String> {
8813        self.prefill_input_rows(PrefillIn::Ids(ids), pos, task_mask)
8814    }
8815
8816    fn prefill_input_rows(
8817        &mut self,
8818        input: PrefillIn<'_>,
8819        pos: usize,
8820        task_mask: Option<&TaskMask>,
8821    ) -> Result<Vec<f32>, String> {
8822        self.mimo_moe_prepare();
8823        let hs = self.hidden_size;
8824        let bk = match input {
8825            PrefillIn::Ids(ids) => ids.len(),
8826            PrefillIn::Hidden(rows) => rows.len() / hs,
8827        };
8828        #[cfg(not(target_os = "macos"))]
8829        if task_mask.is_none()
8830            && !self.o1_active()
8831            && bk > 1
8832            && (self.batch_prefix_prefill()
8833                || (self.verify_exact_moe
8834                    && crate::gpu::enabled_here()
8835                    && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)))
8836        {
8837            let mut hiddens = match input {
8838                PrefillIn::Hidden(rows) => rows.to_vec(),
8839                PrefillIn::Ids(ids) => ids.iter().flat_map(|&id| self.embed_single(id)).collect(),
8840            };
8841            let positions: Vec<usize> = (pos..pos + bk).collect();
8842            let mut run = 0usize;
8843            match self.try_batch_graph_wgpu_prefix(
8844                &mut hiddens,
8845                &positions,
8846                bk,
8847                None,
8848                Some(&mut run),
8849            ) {
8850                crate::gpu::BatchGraphOutcome::Completed => {
8851                    let out = if run < self.num_layers {
8852                        self.prefill_batch_span(
8853                            PrefillIn::Hidden(&hiddens),
8854                            pos,
8855                            None,
8856                            run,
8857                            self.num_layers,
8858                        )
8859                    } else {
8860                        hiddens
8861                    };
8862                    return if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
8863                        Err("MiMo attention graph failed after admission".into())
8864                    } else { Ok(out) };
8865                }
8866                crate::gpu::BatchGraphOutcome::Failed => {
8867                    return Err("batched prefix prefill failed after admission".into());
8868                }
8869                crate::gpu::BatchGraphOutcome::Declined => {
8870                    // Rows an earlier chunk left on the device only.
8871                    #[cfg(feature = "gpu")]
8872                    self.pull_lagging_host_kv(0, self.num_layers, pos);
8873                }
8874            }
8875        }
8876        let out = self.prefill_batch_span(input, pos, task_mask, 0, usize::MAX);
8877        if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
8878            Err("batch tail graph failed after admission".into())
8879        } else { Ok(out) }
8880    }
8881
8882    /// The layer-major batched walk over a layer span [from..upto_excl):
8883    /// the whole prefill machinery (chunk graph, batched attends, GEMM
8884    /// panels) for a PARTIAL stack — the network split's prefill rides
8885    /// the same canon as the local one. Input is token ids (embeds
8886    /// itself, coordinator side) or ready boundary hiddens (worker side).
8887    fn prefill_batch_span(
8888        &mut self,
8889        input: PrefillIn<'_>,
8890        start_pos: usize,
8891        task_mask: Option<&TaskMask>,
8892        from: usize,
8893        upto_excl: usize,
8894    ) -> Vec<f32> {
8895        let hs = self.hidden_size;
8896        let b = match input {
8897            PrefillIn::Ids(ids) => ids.len(),
8898            PrefillIn::Hidden(hb) => hb.len() / hs,
8899        };
8900        let upto_excl = upto_excl.min(self.num_layers);
8901        // The CPU embed is deferred: when the chunk graph takes the run
8902        // from layer 0 it gathers the embeddings on the device instead.
8903        // A hidden input is ready by definition.
8904        let mut h: Vec<f32>;
8905        let mut h_ready;
8906        match input {
8907            PrefillIn::Ids(_) => {
8908                h = vec![0.0; b * hs];
8909                h_ready = false;
8910            }
8911            PrefillIn::Hidden(hb) => {
8912                h = hb.to_vec();
8913                h_ready = true;
8914            }
8915        }
8916        let fill_h = |h: &mut Vec<f32>, me: &Self| {
8917            if let PrefillIn::Ids(ids) = input {
8918                for (bi, &id) in ids.iter().enumerate() {
8919                    let e = me.embed_single(id);
8920                    h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
8921                }
8922                if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
8923                    if let Ok(t) = tp.parse::<usize>() {
8924                        if t >= start_pos && t < start_pos + ids.len() {
8925                            let bi = t - start_pos;
8926                            let row = &h[bi * hs..(bi + 1) * hs];
8927                            let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
8928                            eprintln!(
8929                                "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
8930                                ids[bi],
8931                                row[0],
8932                                row[1],
8933                                ids.len(),
8934                                &ids[..ids.len().min(8)]
8935                            );
8936                        }
8937                    }
8938                }
8939            }
8940        };
8941        let (_nkv, _hd, _rd, eps) = (
8942            self.num_kv_heads,
8943            self.head_dim,
8944            self.rotary_dim,
8945            self.rms_eps,
8946        );
8947        let pool = self.pool.clone();
8948        let norm_style = self.norm_style;
8949        self.mimo_moe_prepare();
8950        let automatic_gpu_prefix = self.automatic_gpu_prefix();
8951
8952        #[cfg(target_os = "macos")]
8953        let mut chunk_skip_until = 0usize;
8954        for li in from..upto_excl {
8955            let _capacity_tail = automatic_gpu_prefix
8956                .filter(|&prefix| {
8957                    li >= prefix && !(self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false))
8958                })
8959                .map(|_| crate::gpu::enter_cpu_scope());
8960            crate::gpu::set_layer(li as i64); // layer-split GPU/CPU
8961            // GPU chunk graph (default-on under CMF_GPU=1): a run of
8962            // consecutive eligible layers for the whole chunk in ONE
8963            // Metal submission — norm, QKV, RoPE with fused mirror
8964            // append, causal attend, O, FFN, hidden device-resident
8965            // across the run. Any refusal falls through to the CPU path.
8966            #[cfg(target_os = "macos")]
8967            if task_mask.is_none() {
8968                if li < chunk_skip_until {
8969                    continue;
8970                }
8971                // Device-side embedding needs a q8_row embedding matrix;
8972                // with any other layout the CPU fills `h` first and the
8973                // graph starts from a ready hidden (refusing the whole
8974                // run over the embedding alone kept q4t models — the
8975                // whole Nanbeige/Bonsai class — on the CPU prefill).
8976                if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
8977                    fill_h(&mut h, self);
8978                    h_ready = true;
8979                }
8980                let ids_for_embed = match input {
8981                    PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
8982                    PrefillIn::Hidden(_) => None,
8983                };
8984                let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
8985                if end > li {
8986                    h_ready = true;
8987                    chunk_skip_until = end;
8988                    // Looped Transformer: the graph stopped at a loop
8989                    // boundary — apply final norm before the next iteration.
8990                    if self.is_loop_end(end - 1) && end < self.num_layers {
8991                        for bi in 0..b {
8992                            let normed = inference::rms_norm(
8993                                &h[bi * hs..(bi + 1) * hs],
8994                                &self.weights.final_norm,
8995                                eps,
8996                                norm_style,
8997                            );
8998                            h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
8999                        }
9000                    }
9001                    continue;
9002                }
9003            }
9004            if !h_ready {
9005                fill_h(&mut h, self);
9006                h_ready = true;
9007            }
9008            if task_mask.is_none() && self.verify_exact_moe {
9009                let positions: Vec<_> = (start_pos..start_pos + b).collect();
9010                match self.mimo_graph_layer_rows(li, &mut h, &positions) {
9011                    crate::gpu::BatchGraphOutcome::Completed => continue,
9012                    crate::gpu::BatchGraphOutcome::Failed => return h,
9013                    crate::gpu::BatchGraphOutcome::Declined => {},
9014                }
9015            }
9016            #[cfg(feature = "gpu")]
9017            self.pull_lagging_host_kv(li, li + 1, start_pos);
9018            let lw = &self.weights.layers[self.phys_layer(li)];
9019            // ── attention ──
9020            match &lw.attn {
9021                AttnKind::Kda(w) => {
9022                    // Projections batched, recurrence sequential.
9023                    let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
9024                    let mut normed = vec![0.0f32; b * hs];
9025                    for bi in 0..b {
9026                        inference::rms_norm_into(
9027                            &h[bi * hs..(bi + 1) * hs],
9028                            &lw.input_norm,
9029                            eps,
9030                            norm_style,
9031                            &mut normed[bi * hs..(bi + 1) * hs],
9032                        );
9033                    }
9034                    let attn = crate::linear_core::kda_forward_batch(
9035                        &normed,
9036                        b,
9037                        w,
9038                        &cfg,
9039                        &mut self.kv_cache.layers[li].linear_state,
9040                        pool.as_deref(),
9041                    );
9042                    for (dst, &a) in h.iter_mut().zip(&attn) {
9043                        *dst += a;
9044                    }
9045                }
9046                AttnKind::LinearGdn(w) => {
9047                    // Projections batched, recurrence sequential.
9048                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
9049                    let mut normed = vec![0.0f32; b * hs];
9050                    for bi in 0..b {
9051                        let r = inference::rms_norm(
9052                            &h[bi * hs..(bi + 1) * hs],
9053                            &lw.input_norm,
9054                            eps,
9055                            norm_style,
9056                        );
9057                        normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9058                    }
9059                    let attn = crate::linear_core::gdn_forward_batch(
9060                        &normed,
9061                        b,
9062                        w,
9063                        &cfg,
9064                        &mut self.kv_cache.layers[li].linear_state,
9065                        pool.as_deref(),
9066                    );
9067                    for (dst, &a) in h.iter_mut().zip(&attn) {
9068                        *dst += a;
9069                    }
9070                }
9071                AttnKind::ShortConv(w) => {
9072                    // Projections batched over the chunk; the conv walks the
9073                    // contiguous positions in order (same ring as decode).
9074                    let cfg = self
9075                        .short_conv_cfg
9076                        .expect("short-conv layer without short_conv_cfg");
9077                    let mut normed = vec![0.0f32; b * hs];
9078                    for bi in 0..b {
9079                        inference::rms_norm_into(
9080                            &h[bi * hs..(bi + 1) * hs],
9081                            &lw.input_norm,
9082                            eps,
9083                            norm_style,
9084                            &mut normed[bi * hs..(bi + 1) * hs],
9085                        );
9086                    }
9087                    let attn = short_conv_forward_batch(
9088                        &normed,
9089                        b,
9090                        w,
9091                        &cfg,
9092                        &mut self.kv_cache.layers[li].linear_state,
9093                        pool.as_deref(),
9094                    );
9095                    for (dst, &a) in h.iter_mut().zip(&attn) {
9096                        *dst += a;
9097                    }
9098                }
9099                AttnKind::Mla(w) => {
9100                    // Per-position prefill (correctness first; latent
9101                    // batching is a later optimization).
9102                    let inv_freq_l = self.layer_inv_freq(li);
9103                    let rs = self.layer_rope_scale(li);
9104                    let mut normed = vec![0.0f32; hs];
9105                    for bi in 0..b {
9106                        inference::rms_norm_into(
9107                            &h[bi * hs..(bi + 1) * hs],
9108                            &lw.input_norm,
9109                            eps,
9110                            norm_style,
9111                            &mut normed,
9112                        );
9113                        let ao = mla_attention(
9114                            w,
9115                            &normed,
9116                            &mut self.kv_cache.layers[li],
9117                            start_pos + bi,
9118                            &inv_freq_l,
9119                            rs,
9120                            eps,
9121                            pool.as_deref(),
9122                        );
9123                        for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
9124                            *dst += a;
9125                        }
9126                    }
9127                }
9128                AttnKind::Full {
9129                    wq,
9130                    wk,
9131                    wv,
9132                    wo,
9133                    q_norm,
9134                    k_norm,
9135                    output_gate,
9136                    softplus_gate,
9137                    bias,
9138                } => {
9139                    // Chunk-GEMM QKV/O; per-position causal attention
9140                    // inside (roadmap §3 P0 — full-attention prefill no
9141                    // longer re-reads the projection weights b times).
9142                    let mut normed = vec![0.0f32; b * hs];
9143                    for bi in 0..b {
9144                        inference::rms_norm_into(
9145                            &h[bi * hs..(bi + 1) * hs],
9146                            &lw.input_norm,
9147                            eps,
9148                            norm_style,
9149                            &mut normed[bi * hs..(bi + 1) * hs],
9150                        );
9151                    }
9152                    let inv_freq_l = self.layer_inv_freq(li);
9153                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
9154                    let cfg = QwenAttnCfg {
9155                        num_heads: self.layer_num_heads(li),
9156                        num_kv_heads: nkv_l,
9157                        head_dim: hd_l,
9158                        hidden_size: hs,
9159                        position: start_pos,
9160                        inv_freq: &inv_freq_l,
9161                        rotary_dim: rd_l,
9162                        scale: self.attn_scale,
9163                        softcap: self.attn_softcap,
9164                        window: self.layer_window(li),
9165                        v_norm: self.attn_v_norm,
9166                        qk_norm_after_rope: self.qk_norm_after_rope,
9167                        q_norm: q_norm.as_deref(),
9168                        k_norm: k_norm.as_deref(),
9169                        output_gate: *output_gate,
9170                        softplus_gate: softplus_gate
9171                            .as_ref()
9172                            .map(|(gate, per_head)| (gate, *per_head)),
9173                        rope_scale: self.layer_rope_scale(li),
9174                        bias: bias
9175                            .as_ref()
9176                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
9177                        rms_eps: eps,
9178                        norm_style,
9179                        pool: pool.as_deref(),
9180                        v_head_dim: self.layer_v_dim(li),
9181                    };
9182                    let mut attn = attention::qwen_attention_batch(
9183                        &normed,
9184                        b,
9185                        wq,
9186                        wk,
9187                        wv,
9188                        wo,
9189                        &mut self.kv_cache.layers[li],
9190                        &cfg,
9191                    );
9192                    if let Some(w) = &lw.attn_out_norm {
9193                        for bi in 0..b {
9194                            inference::rms_norm_into(
9195                                &attn[bi * hs..(bi + 1) * hs],
9196                                w,
9197                                eps,
9198                                norm_style,
9199                                &mut normed[bi * hs..(bi + 1) * hs],
9200                            );
9201                        }
9202                        attn.copy_from_slice(&normed);
9203                    }
9204                    for (dst, &a) in h.iter_mut().zip(&attn) {
9205                        *dst += a;
9206                    }
9207                }
9208                AttnKind::Bounded(w) => {
9209                    // Chunk-GEMM projections, the bounded operator per
9210                    // position over ring + chunk — never a growing KV.
9211                    let mut normed = vec![0.0f32; b * hs];
9212                    for bi in 0..b {
9213                        inference::rms_norm_into(
9214                            &h[bi * hs..(bi + 1) * hs],
9215                            &lw.input_norm,
9216                            eps,
9217                            norm_style,
9218                            &mut normed[bi * hs..(bi + 1) * hs],
9219                        );
9220                    }
9221                    let rope = self
9222                        .bounded_rope
9223                        .clone()
9224                        .expect("bounded layer without an installed rotation table");
9225                    let cfg = crate::bounded::BoundedAttnCfg {
9226                        num_heads: self.num_heads,
9227                        num_kv_heads: self.num_kv_heads,
9228                        head_dim: self.head_dim,
9229                        hidden_size: hs,
9230                        scale: self.attn_scale,
9231                        rope: &rope,
9232                        pool: pool.as_deref(),
9233                    };
9234                    let mut attn = crate::bounded::bounded_attention_batch(
9235                        &normed,
9236                        b,
9237                        w,
9238                        &mut self.kv_cache.layers[li],
9239                        &cfg,
9240                    );
9241                    if let Some(wn) = &lw.attn_out_norm {
9242                        for bi in 0..b {
9243                            inference::rms_norm_into(
9244                                &attn[bi * hs..(bi + 1) * hs],
9245                                wn,
9246                                eps,
9247                                norm_style,
9248                                &mut normed[bi * hs..(bi + 1) * hs],
9249                            );
9250                        }
9251                        attn.copy_from_slice(&normed);
9252                    }
9253                    for (dst, &a) in h.iter_mut().zip(&attn) {
9254                        *dst += a;
9255                    }
9256                    attention::recycle_buf(&mut attn);
9257                }
9258                AttnKind::Linear(w) => {
9259                    for bi in 0..b {
9260                        let normed = inference::rms_norm(
9261                            &h[bi * hs..(bi + 1) * hs],
9262                            &lw.input_norm,
9263                            eps,
9264                            norm_style,
9265                        );
9266                        vmf_phase_forward(
9267                            &normed,
9268                            w,
9269                            &self.vmf_cfg.expect("linear layer without vmf_cfg"),
9270                            &mut self.kv_cache.layers[li].linear_state,
9271                            pool.as_deref(),
9272                        )
9273                        .iter()
9274                        .enumerate()
9275                        .for_each(|(i, &a)| h[bi * hs + i] += a);
9276                    }
9277                }
9278            }
9279
9280            // ── FFN batched ──
9281            let lw = &self.weights.layers[self.phys_layer(li)];
9282            let mut post = vec![0.0f32; b * hs];
9283            for bi in 0..b {
9284                let r =
9285                    inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
9286                post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9287            }
9288            // A restrictive per-visit FFN row lands on the activations
9289            // inside the dense arm; an all-open row costs nothing.
9290            let mask_row = task_mask
9291                .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
9292                .and_then(|m| m.ffn_masks.get(li))
9293                .map(|v| v.as_slice());
9294            let mut ffn = match &lw.ffn {
9295                FfnKind::Dense(d) if !d.segs.is_empty() => {
9296                    tube_ffn(d, &post, b, pool.as_deref(), mask_row)
9297                }
9298                FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
9299                FfnKind::Moe(m) if self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false) => {
9300                    moe_ffn_banked_rows(&mut self.mimo_moe, li, m, &post, b, hs, pool.as_deref())
9301                }
9302                FfnKind::Moe(m) if self.verify_exact_moe => {
9303                    moe_ffn_rows_exact(m, &post, b, hs, pool.as_deref())
9304                }
9305                // Keep prompt expert panels off the projection arena and
9306                // use their routes to prime the model-wide bank.
9307                FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => {
9308                    let before = m.stats.borrow().clone();
9309                    let out = crate::gpu::cpu_scope(|| {
9310                        moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None)
9311                    });
9312                    self.mimo_moe.prime(li, m, &before);
9313                    out
9314                }
9315                FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
9316                // Dual-branch layers run per position (the expert branch
9317                // reads the raw residual — nothing to batch yet).
9318                FfnKind::DenseMoe(dm) => {
9319                    let mut out = vec![0.0f32; b * hs];
9320                    for bi in 0..b {
9321                        let r = dense_moe_ffn(
9322                            dm,
9323                            &post[bi * hs..(bi + 1) * hs],
9324                            &h[bi * hs..(bi + 1) * hs],
9325                            eps,
9326                            norm_style,
9327                            pool.as_deref(),
9328                        );
9329                        out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9330                    }
9331                    out
9332                }
9333            };
9334            if let Some(w) = &lw.ffn_out_norm {
9335                for bi in 0..b {
9336                    inference::rms_norm_into(
9337                        &ffn[bi * hs..(bi + 1) * hs],
9338                        w,
9339                        eps,
9340                        norm_style,
9341                        &mut post[bi * hs..(bi + 1) * hs],
9342                    );
9343                }
9344                ffn.copy_from_slice(&post);
9345            }
9346            for (dst, &f) in h.iter_mut().zip(&ffn) {
9347                *dst += f;
9348            }
9349            if let Some(sc) = lw.layer_scale {
9350                for v in h.iter_mut() {
9351                    *v *= sc;
9352                }
9353            }
9354            // CMF_LAYER_DUMP: every position's hidden after layer li.
9355            if self.layer_dump.is_some() {
9356                for bi in 0..b {
9357                    self.dump_layer_row(start_pos + bi, li, &h[bi * hs..(bi + 1) * hs]);
9358                }
9359            }
9360            if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9361                if let Ok(t) = tp.parse::<usize>() {
9362                    if t >= start_pos && t < start_pos + b {
9363                        let bi = t - start_pos;
9364                        let row = &h[bi * hs..(bi + 1) * hs];
9365                        let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9366                        eprintln!(
9367                            "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
9368                            row[0], row[1]
9369                        );
9370                    }
9371                }
9372            }
9373            // CMF_DEBUG_LAYERS=1: per-layer hidden-state health of the
9374            // LAST prompt position — the knife for "which layer type
9375            // breaks first" on a new architecture.
9376            if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
9377                let row = &h[(b - 1) * hs..b * hs];
9378                let rms =
9379                    (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
9380                let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
9381                eprintln!(
9382                    "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
9383                    match &self.weights.layers[self.phys_layer(li)].attn {
9384                        AttnKind::LinearGdn(_) => "gdn",
9385                        AttnKind::Linear(_) => "vmf",
9386                        AttnKind::ShortConv(_) => "conv",
9387                        _ => "attn",
9388                    },
9389                    match &lw.ffn {
9390                        FfnKind::Moe(_) => "moe",
9391                        FfnKind::Dense(_) => "dense",
9392                        FfnKind::DenseMoe(_) => "dense+moe",
9393                    },
9394                );
9395            }
9396            // Looped Transformer: apply final norm at the end of each loop iteration.
9397            if self.is_loop_end(li) && li + 1 < self.num_layers {
9398                for bi in 0..b {
9399                    let normed = inference::rms_norm(
9400                        &h[bi * hs..(bi + 1) * hs],
9401                        &self.weights.final_norm,
9402                        eps,
9403                        norm_style,
9404                    );
9405                    h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9406                }
9407            }
9408            if std::env::var("CMF_TRACE_H").is_ok() {
9409                let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
9410                let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
9411                eprintln!(
9412                    "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
9413                    lw.layer_scale
9414                );
9415            }
9416        }
9417        crate::gpu::set_layer(-1); // lm_head/final ops outside layer-split
9418        // A batched span owns a complete set of positions. Publish any
9419        // collecting→sealed transition only after every layer has finished;
9420        // callers that cross into serial/device work must see the new epoch
9421        // before this function returns.
9422        self.o1_progress();
9423        h
9424    }
9425
9426    /// Embed a single token.
9427    fn embed_single(&self, id: u32) -> Vec<f32> {
9428        let mut out = vec![0.0f32; self.hidden_size];
9429        if (id as usize) < self.weights.embed_tokens.rows() {
9430            self.weights.embed_tokens.row_f32(id as usize, &mut out);
9431        }
9432        if self.embed_multiplier != 1.0 {
9433            for v in out.iter_mut() {
9434                *v *= self.embed_multiplier;
9435            }
9436        }
9437        // DeepSeek-V4's hash layers route by TOKEN ID, so the id has to
9438        // reach the forward. It rides in slot 0 (the forward re-reads the
9439        // real embedding itself from the table).
9440        if self.dsv4.is_some()
9441            || self.dsv41.is_some()
9442            || self.qwen4_exp.is_some()
9443        {
9444            let mut v = vec![0.0f32; self.hidden_size.max(1)];
9445            v[0] = id as f32;
9446            return v;
9447        }
9448        // Gemma-3n: the per-layer-embedding half needs the token ID, so
9449        // it rides appended to the embedding; the g3n forward splits it.
9450        if let Some(b) = &self.g3n {
9451            return b.0.extend_embedding(id, &out, self.pool.as_deref());
9452        }
9453        out
9454    }
9455
9456    /// A run of consecutive prefill layers on the GPU for the whole
9457    /// chunk (default-on under CMF_GPU=1; CMF_GPU_CHUNK=0 disables).
9458    /// Eligibility per layer: q8_row weights, plain full attention
9459    /// (no output gate), F32 KV, no o1/masks/gemma extras. Returns the
9460    /// first layer index NOT processed (== `li0` when the run is empty).
9461    #[cfg(target_os = "macos")]
9462    fn chunk_run_gpu(
9463        &mut self,
9464        li0: usize,
9465        h: &mut [f32],
9466        b: usize,
9467        pos0: usize,
9468        embed_ids: Option<&[u32]>,
9469        cap: usize,
9470    ) -> usize {
9471        // (The old streaming attend needed a depth bound at ~1k; the
9472        // GEMM attention scales like the CPU path and lifted it.)
9473        // CMF_GPU_CHUNK=0 disables the graph.
9474        if !crate::gpu::enabled_here()
9475            || std::env::var("CMF_GPU_CHUNK")
9476                .map(|v| v == "0")
9477                .unwrap_or(false)
9478            || b < 32
9479            || self.swa.is_some()
9480            || self.global_attn.is_some()
9481            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
9482            || self.graph_attn_decline_reason().is_some()
9483            // Collection owns the exact Q trace and boundary conversion;
9484            // this chunk graph appends dense KV without feeding that trace.
9485            || self.o1_active()
9486            || self.attn_v_norm
9487            || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
9488        {
9489            return li0;
9490        }
9491        let Some(model) = self.model.clone() else {
9492            return li0;
9493        };
9494        let inv_freq = self.inv_freq.clone();
9495        let (nh, nkv, hd, hs) = (
9496            self.num_heads,
9497            self.num_kv_heads,
9498            self.head_dim,
9499            self.hidden_size,
9500        );
9501        // Collect the longest run of consecutive eligible layers.
9502        // Looped Transformer: stop at the loop boundary so the CPU can
9503        // apply loop_final_norm between iterations.
9504        let loop_end = if self.loop_final_norm {
9505            ((li0 / self.physical_layers) + 1) * self.physical_layers
9506        } else {
9507            self.num_layers
9508        };
9509        let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
9510        let mut stored_at: Vec<usize> = Vec::new();
9511        for li in li0..self.num_layers.min(loop_end).min(cap) {
9512            let lw = &self.weights.layers[self.phys_layer(li)];
9513            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
9514                break;
9515            }
9516            let AttnKind::Full {
9517                wq,
9518                wk,
9519                wv,
9520                wo,
9521                q_norm,
9522                k_norm,
9523                output_gate: false,
9524                softplus_gate: None,
9525                bias,
9526            } = &lw.attn
9527            else {
9528                break;
9529            };
9530            let FfnKind::Dense(d) = &lw.ffn else { break };
9531            if d.act != Act::Silu || !d.segs.is_empty() {
9532                break;
9533            }
9534            // q8_row (row_scale populated), or q4_tiled / q4tp (row_scale
9535            // empty — their scales are in the payload). Mixing across the
9536            // seven projections of one layer is fine; the encoder branches
9537            // per weight on the tensor's dtype. Anything else refuses.
9538            fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
9539                t.q8_row_parts()
9540                    .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9541                    .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9542            }
9543            let parts = (
9544                cw(wq),
9545                cw(wk),
9546                cw(wv),
9547                cw(wo),
9548                cw(&d.gate_proj),
9549                cw(&d.up_proj),
9550                cw(&d.down_proj),
9551            );
9552            let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
9553            else {
9554                break;
9555            };
9556            let layer = &self.kv_cache.layers[li];
9557            if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
9558                break;
9559            }
9560            stored_at.push(layer.head_len(0));
9561            layers.push(crate::gpu_metal::ChunkLayer {
9562                model: &model,
9563                kv_id: self.graph_kv_id,
9564                layer: li,
9565                wq: pq,
9566                wk: pk,
9567                wv: pv,
9568                wo: po,
9569                gate: pg,
9570                up: pu,
9571                down: pd,
9572                input_norm: &lw.input_norm,
9573                post_norm: &lw.post_norm,
9574                bias: bias
9575                    .as_ref()
9576                    .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
9577                q_norm: q_norm.as_deref(),
9578                k_norm: k_norm.as_deref(),
9579                inv_freq: &inv_freq,
9580                rd: self.rotary_dim,
9581                nh,
9582                nkv,
9583                hd,
9584                hs,
9585                inter: d.gate_proj.rows(),
9586                gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
9587                late_qk_norm: self.qk_norm_after_rope,
9588                eps: self.rms_eps as f32,
9589            });
9590        }
9591        if layers.is_empty() {
9592            return li0;
9593        }
9594        let row = nkv * hd;
9595        let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
9596            .iter()
9597            .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
9598            .collect();
9599        let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
9600        for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
9601            let li = layers[i].layer;
9602            let layer = &self.kv_cache.layers[li];
9603            io.push(crate::gpu_metal::ChunkIo {
9604                cpu_stored: stored_at[i],
9605                cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
9606                cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
9607                out_k: ok,
9608                out_v: ov,
9609                imp: oi,
9610            });
9611        }
9612        let n_run = layers.len();
9613        let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
9614        // Device-side embedding when the run starts the model and the
9615        // embedding matrix is q8_row-mapped.
9616        let ep = embed_ids.and_then(|ids| {
9617            self.weights
9618                .embed_tokens
9619                .q8_row_parts()
9620                .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
9621                    idx,
9622                    rows,
9623                    row_scale: rs,
9624                    ids,
9625                    mult: self.embed_multiplier,
9626                })
9627        });
9628        if embed_ids.is_some() && ep.is_none() {
9629            return li0;
9630        }
9631        if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
9632            return li0;
9633        }
9634        drop(io);
9635        drop(layers);
9636        // CPU caches stay the owners of record: append the chunk rows
9637        // and bank the importance masses per layer.
9638        for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
9639            let li = li0 + i;
9640            let layer = &mut self.kv_cache.layers[li];
9641            for bi in 0..b {
9642                layer.append(
9643                    &ok[bi * row..(bi + 1) * row],
9644                    &ov[bi * row..(bi + 1) * row],
9645                    &[],
9646                );
9647            }
9648            layer.accumulate_imp(oi);
9649        }
9650        last
9651    }
9652
9653    /// Is layer `li` a sliding-window (local-RoPE) layer? Gemma-3:
9654    /// every `pattern`-th layer is global, the rest are local.
9655    fn layer_is_local(&self, li: usize) -> bool {
9656        if let Some(layers) = &self.sliding_layers {
9657            return layers.get(li).copied().unwrap_or(false);
9658        }
9659        match self.swa {
9660            Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
9661            None => false,
9662        }
9663    }
9664
9665    /// The RoPE table for layer `li` (local layers may have their own;
9666    /// Gemma-4 global layers use the proportional padded table).
9667    fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
9668        if self.layer_is_local(li) {
9669            if let Some(f) = &self.inv_freq_local {
9670                return f.clone();
9671            }
9672        } else if let Some(f) = &self.inv_freq_global {
9673            return f.clone();
9674        }
9675        self.inv_freq.clone()
9676    }
9677
9678    /// The attend window for layer `li` (None = full context).
9679    fn layer_window(&self, li: usize) -> Option<usize> {
9680        self.swa
9681            .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
9682    }
9683
9684    fn layer_num_heads(&self, li: usize) -> usize {
9685        self.attention_heads_per_layer
9686            .as_ref()
9687            .and_then(|v| v.get(li).copied())
9688            .unwrap_or(self.num_heads)
9689    }
9690
9691    fn layer_rope_scale(&self, li: usize) -> f32 {
9692        if self.layer_is_local(li) {
9693            self.rope_scale_local
9694        } else {
9695            self.rope_scale
9696        }
9697    }
9698
9699    /// Attention geometry of layer `li`: (num_kv_heads, head_dim,
9700    /// rotary_dim). Gemma-4 global layers override all three.
9701    fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
9702        if !self.layer_is_local(li) {
9703            if let Some((ghd, gkv)) = self.global_attn {
9704                return (gkv, ghd, ghd);
9705            }
9706        }
9707        (
9708            self.layer_num_kv_heads(li),
9709            self.head_dim,
9710            if self.layer_is_local(li) {
9711                self.rotary_dim_local.unwrap_or(self.rotary_dim)
9712            } else {
9713                self.rotary_dim
9714            },
9715        )
9716    }
9717
9718    /// KV heads of layer `li` (virtual index): the per-layer count when the
9719    /// model has one (MiMo-V2), else the uniform `num_kv_heads`.
9720    fn layer_num_kv_heads(&self, li: usize) -> usize {
9721        self.kv_heads_per_layer
9722            .as_ref()
9723            .and_then(|v| v.get(self.phys_layer(li)).copied())
9724            .unwrap_or(self.num_kv_heads)
9725    }
9726
9727    /// V head width of layer `li` (≤ its head_dim).
9728    fn layer_v_dim(&self, li: usize) -> usize {
9729        let (_, hd, _) = self.layer_geom(li);
9730        self.v_head_dim.unwrap_or(hd).min(hd)
9731    }
9732
9733    /// Install a per-layer KV geometry: KV heads per PHYSICAL layer and/or
9734    /// a V head width narrower than `head_dim` (MiMo-V2). Validates it and
9735    /// reshapes the caches of every layer whose KV head count differs from
9736    /// `num_kv_heads`. The loader and the tests share this one path, so a
9737    /// hand-built pipeline cannot hold a geometry the loader would refuse.
9738    /// Call before the first forward (it drops cached rows of reshaped
9739    /// layers). Refuses combinations whose paths would read it wrong:
9740    /// Gemma-4 global layers and MLA carry their own geometry.
9741    pub fn set_attn_geometry(
9742        &mut self,
9743        kv_heads_per_layer: Option<Vec<usize>>,
9744        v_head_dim: Option<usize>,
9745    ) -> Result<(), String> {
9746        if kv_heads_per_layer.is_some() || v_head_dim.is_some() {
9747            if self.global_attn.is_some() {
9748                return Err(
9749                    "per-layer KV heads / v_head_dim cannot combine with Gemma-4 global \
9750                     attention geometry"
9751                        .into(),
9752                );
9753            }
9754            if self
9755                .weights
9756                .layers
9757                .iter()
9758                .any(|lw| matches!(lw.attn, AttnKind::Mla(_)))
9759            {
9760                return Err("per-layer KV heads / v_head_dim cannot combine with MLA".into());
9761            }
9762        }
9763        if let Some(vd) = v_head_dim {
9764            if vd == 0 || vd > self.head_dim {
9765                return Err(format!(
9766                    "v_head_dim {vd} must be in 1..={} (head_dim)",
9767                    self.head_dim
9768                ));
9769            }
9770        }
9771        if let Some(v) = &kv_heads_per_layer {
9772            if v.len() != self.physical_layers {
9773                return Err(format!(
9774                    "kv_heads_per_layer has {} entries, expected {} layers",
9775                    v.len(),
9776                    self.physical_layers
9777                ));
9778            }
9779            for (li, &nkv) in v.iter().enumerate() {
9780                let is_attn = matches!(
9781                    self.weights.layers.get(li).map(|lw| &lw.attn),
9782                    Some(AttnKind::Full { .. }) | None
9783                );
9784                if !is_attn {
9785                    continue;
9786                }
9787                let nh = self
9788                    .attention_heads_per_layer
9789                    .as_ref()
9790                    .and_then(|h| h.get(li).copied())
9791                    .unwrap_or(self.num_heads);
9792                if nkv == 0 || nh % nkv != 0 {
9793                    return Err(format!(
9794                        "layer {li}: {nkv} KV heads must be nonzero and divide {nh} Q heads"
9795                    ));
9796                }
9797            }
9798        }
9799        self.kv_heads_per_layer = kv_heads_per_layer;
9800        self.v_head_dim = v_head_dim.filter(|&vd| vd != self.head_dim);
9801        if self.kv_heads_per_layer.is_some() {
9802            for li in 0..self.kv_cache.layers.len() {
9803                let full = matches!(
9804                    self.weights
9805                        .layers
9806                        .get(self.phys_layer(li))
9807                        .map(|lw| &lw.attn),
9808                    Some(AttnKind::Full { .. })
9809                );
9810                let nkv = self.layer_num_kv_heads(li);
9811                let cache = &self.kv_cache.layers[li];
9812                if full && (cache.num_kv_heads != nkv || cache.head_dim != self.head_dim) {
9813                    let sinks = cache.sinks.clone();
9814                    self.kv_cache.layers[li] =
9815                        crate::kv_cache::LayerKvCache::new(nkv, self.head_dim);
9816                    self.kv_cache.layers[li].sinks = sinks;
9817                }
9818            }
9819        }
9820        Ok(())
9821    }
9822
9823    /// Attach learned attention-sink logits (one per Q head) to PHYSICAL
9824    /// layer `phys` — every virtual layer that runs it. The loader calls
9825    /// this for each `model.layers.N.self_attn.sinks` tensor.
9826    pub fn set_layer_sinks(&mut self, phys: usize, sinks: Vec<f32>) -> Result<(), String> {
9827        let Some(lw) = self.weights.layers.get(phys) else {
9828            return Err(format!("sinks for layer {phys}: no such layer"));
9829        };
9830        if !matches!(lw.attn, AttnKind::Full { .. }) {
9831            return Err(format!(
9832                "sinks for layer {phys}: only softmax (Full) attention layers take sinks"
9833            ));
9834        }
9835        let nh = self
9836            .attention_heads_per_layer
9837            .as_ref()
9838            .and_then(|h| h.get(phys).copied())
9839            .unwrap_or(self.num_heads);
9840        if sinks.len() != nh {
9841            return Err(format!(
9842                "sinks for layer {phys}: {} values, expected one per Q head ({nh})",
9843                sinks.len()
9844            ));
9845        }
9846        if let Some(bad) = sinks.iter().find(|v| !v.is_finite()) {
9847            return Err(format!("sinks for layer {phys}: non-finite value {bad}"));
9848        }
9849        for li in 0..self.kv_cache.layers.len() {
9850            if self.phys_layer(li) == phys {
9851                self.kv_cache.layers[li].sinks = Some(sinks.clone());
9852            }
9853        }
9854        Ok(())
9855    }
9856
9857    /// Why the GPU attention graphs cannot serve this model, if they
9858    /// cannot: the wgpu whole-token and batched graphs, the greedy
9859    /// multi-burst, the q1 attention dropin and the Metal block/chunk/rows
9860    /// graphs all assume ONE (num_kv_heads, head_dim) geometry, V heads as
9861    /// wide as K, a single RoPE table, full-context attention and a plain
9862    /// softmax. A model outside that contract runs on the CPU layer walk
9863    /// (and the per-op GPU matvecs) until a graph learns it — never on a
9864    /// graph that would read it wrong. None = no attention-level reason
9865    /// (the graph builders still check weights and layer kinds).
9866    pub fn graph_attn_decline_reason(&self) -> Option<&'static str> {
9867        if self.kv_heads_per_layer.is_some() {
9868            return Some("per-layer KV head counts");
9869        }
9870        if self.v_head_dim.is_some_and(|vd| vd != self.head_dim) {
9871            return Some("V heads narrower than Q/K heads");
9872        }
9873        if self.kv_cache.layers.iter().any(|l| l.sinks.is_some()) {
9874            return Some("learned attention sinks");
9875        }
9876        if self.swa.is_some() || self.sliding_layers.is_some() {
9877            return Some("sliding-window layers");
9878        }
9879        None
9880    }
9881
9882    /// Why the WGPU graphs (whole-token, batched prefill, greedy burst)
9883    /// cannot run this model's attention, if they cannot. Per-layer KV
9884    /// heads, V narrower than K, learned sinks and sliding windows ride
9885    /// their per-layer geometry (`GraphAttnGeom`, the ATTEND_X kernels);
9886    /// what that geometry does not express keeps the decline, by name.
9887    /// None for every model with one attention geometry.
9888    pub fn wgpu_graph_attn_decline(&self) -> Option<&'static str> {
9889        self.graph_attn_decline_reason()?;
9890        if self.global_attn.is_some() {
9891            return Some("per-layer head width (Gemma-4 global layers) with per-layer geometry");
9892        }
9893        if self.attention_heads_per_layer.is_some() {
9894            return Some("per-layer Q head counts with per-layer geometry");
9895        }
9896        if self.attn_v_norm {
9897            return Some("V norm with per-layer geometry");
9898        }
9899        if (0..self.num_layers).any(|li| self.layer_rope_scale(li) != 1.0) {
9900            return Some("scaled RoPE positions with per-layer geometry");
9901        }
9902        if self.weights.layers.iter().any(|lw| {
9903            lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some()
9904        }) {
9905            return Some("sandwich norms / layer scale with per-layer geometry");
9906        }
9907        if self.weights.layers.iter().any(|lw| {
9908            matches!(
9909                &lw.attn,
9910                AttnKind::Full {
9911                    output_gate: true,
9912                    ..
9913                }
9914            )
9915        }) && self.v_head_dim.is_some()
9916        {
9917            return Some("gated attention with V narrower than K");
9918        }
9919        if (0..self.num_layers).any(|li| {
9920            self.layer_is_local(li)
9921                && self.inv_freq_local.is_none()
9922                && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
9923        }) {
9924            return Some("local rotary width without a local RoPE table");
9925        }
9926        None
9927    }
9928
9929    /// The wgpu graphs' attention geometry for layer `li` (virtual index):
9930    /// Some only for a model whose layers do not share one (MiMo-V2) — KV
9931    /// heads, V width, rotary width and RoPE table, window and sinks of
9932    /// THIS layer, exactly what the CPU attention reads for it.
9933    fn graph_attn_geom(&self, li: usize) -> Option<crate::gpu::GraphAttnGeom<'_>> {
9934        self.graph_attn_decline_reason()?;
9935        let (nkv, _hd, rd) = self.layer_geom(li);
9936        let invf: &[f32] = if self.layer_is_local(li) {
9937            match &self.inv_freq_local {
9938                Some(f) => f.as_slice(),
9939                None => self.inv_freq.as_slice(),
9940            }
9941        } else {
9942            match &self.inv_freq_global {
9943                Some(f) => f.as_slice(),
9944                None => self.inv_freq.as_slice(),
9945            }
9946        };
9947        Some(crate::gpu::GraphAttnGeom {
9948            nkv,
9949            dv: self.layer_v_dim(li),
9950            rd,
9951            invf,
9952            window: self.layer_window(li),
9953            sink: self.kv_cache.layers[li].sinks.as_deref(),
9954        })
9955    }
9956
9957    /// Bring the host KV cache of every Full-attention layer in
9958    /// `[from, upto)` up to `position` rows from the wgpu mirrors, where a
9959    /// device graph advanced a layer that the host is about to run: a
9960    /// device prefix that shrank since the prompt (or a batched prefill
9961    /// prefix longer than the decode one). A layer whose mirror does not
9962    /// hold the missing rows is left alone. Rows a sliding layer's ring
9963    /// no longer holds come back as zeros — outside every window that
9964    /// will read them.
9965    #[cfg(feature = "gpu")]
9966    fn pull_lagging_host_kv(&mut self, from: usize, upto: usize, position: usize) {
9967        let kv_id = self.graph_kv_id;
9968        for li in from..upto.min(self.num_layers) {
9969            if !matches!(
9970                self.weights.layers[self.phys_layer(li)].attn,
9971                AttnKind::Full { .. }
9972            ) {
9973                continue;
9974            }
9975            let host = self.kv_cache.layers[li].seq_len;
9976            if host >= position {
9977                continue;
9978            }
9979            let Some(dev) = crate::gpu::graph_kv_stored(kv_id, li) else {
9980                continue;
9981            };
9982            let to = dev.min(position);
9983            if to <= host {
9984                continue;
9985            }
9986            let (nkv, hd) = {
9987                let c = &self.kv_cache.layers[li];
9988                (c.num_kv_heads, c.head_dim)
9989            };
9990            let Some((k, v, first_valid)) =
9991                crate::gpu::graph_kv_pull_host(kv_id, li, host, to, nkv, hd)
9992            else {
9993                continue;
9994            };
9995            // A sliding layer only ever reads its last `window` rows; a
9996            // full-context layer needs every row it did not have.
9997            let need_from = match self.layer_window(li) {
9998                Some(w) => host.max((position + 1).saturating_sub(w)),
9999                None => host,
10000            };
10001            if first_valid > need_from {
10002                tracing::warn!(
10003                    "layer {li}: device KV rows {host}..{to} no longer resident \
10004                     (from {first_valid}); host attention will miss them"
10005                );
10006            }
10007            let row = nkv * hd;
10008            let cache = &mut self.kv_cache.layers[li];
10009            for p in 0..to - host {
10010                cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
10011            }
10012        }
10013    }
10014
10015    /// Log (once per graph site and pipeline) that `site` declined for
10016    /// `reason`. The lines are kept so a caller or a test can read them.
10017    fn note_graph_decline(&self, site: &'static str, reason: &'static str) {
10018        let mut seen = self.graph_declines.borrow_mut();
10019        if !seen.iter().any(|&(s, r)| s == site && r == reason) {
10020            tracing::warn!("{site} declined: {reason} (CPU attention path)");
10021            seen.push((site, reason));
10022        }
10023    }
10024
10025    /// The GPU-graph declines this pipeline has logged so far, as the
10026    /// logged lines.
10027    pub fn graph_declines(&self) -> Vec<String> {
10028        self.graph_declines
10029            .borrow()
10030            .iter()
10031            .map(|(site, reason)| format!("{site} declined: {reason} (CPU attention path)"))
10032            .collect()
10033    }
10034
10035    /// Does layer `li` have the plain attention geometry the historical
10036    /// head-masked f32 path (`multi_head_attention`) assumes — pipeline-wide
10037    /// KV heads / head_dim / RoPE table, full context, no sink, V as wide
10038    /// as K? Anything else runs the dense `qwen_attention` instead.
10039    /// `CMF_LAYER_DUMP` writer (see `Pipeline::layer_dump`): one position's
10040    /// hidden after layer `li` as raw little-endian f32 into
10041    /// `<dir>/p{pos:06}_l{li:02}.f32`. A failed write is reported once and
10042    /// never stops the forward — the dump is a diagnostic.
10043    fn dump_layer_row(&self, pos: usize, li: usize, row: &[f32]) {
10044        let Some(dir) = &self.layer_dump else {
10045            return;
10046        };
10047        let mut bytes = Vec::with_capacity(row.len() * 4);
10048        for v in row {
10049            bytes.extend_from_slice(&v.to_le_bytes());
10050        }
10051        let path = dir.join(format!("p{pos:06}_l{li:02}.f32"));
10052        if let Err(e) = std::fs::create_dir_all(dir).and_then(|_| std::fs::write(&path, &bytes)) {
10053            use std::sync::atomic::{AtomicBool, Ordering};
10054            static SAID: AtomicBool = AtomicBool::new(false);
10055            if !SAID.swap(true, Ordering::Relaxed) {
10056                tracing::error!("CMF_LAYER_DUMP: cannot write {}: {e}", path.display());
10057            }
10058        }
10059    }
10060
10061    /// Decide the MiMo-V2 expert placement once (`crate::mimo_moe`). Any
10062    /// other model turns the slot off on the first call.
10063    fn mimo_moe_prepare(&mut self) {
10064        if !self.mimo_moe.is_undecided() {
10065            return;
10066        }
10067        let slot = {
10068            let layers: Vec<(usize, &MoeFfn)> = (0..self.num_layers)
10069                .filter_map(
10070                    |li| match &self.weights.layers.get(self.phys_layer(li))?.ffn {
10071                        FfnKind::Moe(m) => Some((li, m)),
10072                        _ => None,
10073                    },
10074                )
10075                .collect();
10076            // One bank lives on one device: an in-process multi-GPU split
10077            // keeps the whole-layer path.
10078            if layers.is_empty()
10079                || self.physical_layers != self.num_layers
10080                || self.gpu_plan.is_some()
10081            {
10082                crate::mimo_moe::Slot::Off
10083            } else {
10084                // Whether a whole-token graph could run this model's layers
10085                // (then a whole-layer prefix is one submit, not per-layer
10086                // fences).
10087                let graph_prefix =
10088                    self.wgpu_graph_attn_decline().is_none() && crate::gpu::wgpu_graph_default();
10089                crate::mimo_moe::Slot::decide(&layers, self.num_layers, graph_prefix)
10090            }
10091        };
10092        self.mimo_moe = slot;
10093    }
10094
10095    #[cfg(test)]
10096    pub(crate) fn test_graph_kv_id(&self) -> u64 {
10097        self.graph_kv_id
10098    }
10099
10100    /// Dynamic MiMo layer: one device attention graph, followed by a bank
10101    /// frame. Both decode and short verification use this same attention
10102    /// path and absolute layer key; the host KV may intentionally lag.
10103    pub(crate) fn mimo_graph_layer_rows(
10104        &mut self,
10105        li: usize,
10106        h: &mut [f32],
10107        positions: &[usize],
10108    ) -> crate::gpu::BatchGraphOutcome {
10109        use crate::gpu::BatchGraphOutcome as Out;
10110        let b = positions.len();
10111        if !(1..=4).contains(&b)
10112            || h.len() != b * self.hidden_size
10113            || !self.mimo_moe.is_dynamic(li, true)
10114            || !crate::gpu::enabled_here()
10115            || !crate::gpu::wgpu_active()
10116            || self.o1_active()
10117            || self.physical_layers != self.num_layers
10118            // The pair-fusion diagnostic (and an explicit graph-off run)
10119            // rewinds only host KV. A hidden singleton attention graph here
10120            // would leave device mirrors ahead of the next host position.
10121            || std::env::var("CMF_GPU_WGPU_GRAPH").as_deref() == Ok("0")
10122            || std::env::var("CMF_MIMO_ATTN_GRAPH").as_deref() == Ok("0")
10123            || self.wgpu_graph_attn_decline().is_some()
10124        {
10125            return Out::Declined;
10126        }
10127        let attn_started = std::time::Instant::now();
10128        let outcome = {
10129            let lw = &self.weights.layers[li];
10130            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
10131                return Out::Declined;
10132            }
10133            let FfnKind::Moe(m) = &lw.ffn else {
10134                return Out::Declined;
10135            };
10136            let AttnKind::Full {
10137                wq,
10138                wk,
10139                wv,
10140                wo,
10141                q_norm,
10142                k_norm,
10143                output_gate,
10144                softplus_gate,
10145                bias,
10146            } = &lw.attn
10147            else {
10148                return Out::Declined;
10149            };
10150            if *output_gate || softplus_gate.is_some() {
10151                return Out::Declined;
10152            }
10153            let Some(model) = wq.graph_weight().map(|(m, ..)| m.clone()).or_else(|| {
10154                m.experts
10155                    .first()?
10156                    .gate_proj
10157                    .mapped_q4tp()
10158                    .map(|(m, _)| m.clone())
10159            }) else {
10160                return Out::Declined;
10161            };
10162            fn gw<'a>(
10163                t: &'a QTensor,
10164                owner: &std::sync::Arc<cortiq_core::CmfModel>,
10165            ) -> Option<crate::gpu::GraphW<'a>> {
10166                if let Some((m, idx, kind, rs)) = t.graph_weight() {
10167                    if m.uid() != owner.uid() || t.has_prism_contract() {
10168                        return None;
10169                    }
10170                    return Some(crate::gpu::GraphW {
10171                        idx,
10172                        kind,
10173                        row_scale: rs,
10174                        data: &[],
10175                        prism: crate::gpu::GraphPrismOp::None,
10176                        affine: false,
10177                    });
10178                }
10179                t.as_f32().map(|data| crate::gpu::GraphW {
10180                    idx: 0,
10181                    kind: 4,
10182                    row_scale: &[],
10183                    data,
10184                    prism: crate::gpu::GraphPrismOp::None,
10185                    affine: false,
10186                })
10187            }
10188            let (Some(q), Some(k), Some(v), Some(o)) = (
10189                gw(wq, &model),
10190                gw(wk, &model),
10191                gw(wv, &model),
10192                gw(wo, &model),
10193            ) else {
10194                return Out::Declined;
10195            };
10196            let layer = crate::gpu::GraphLayer {
10197                input_norm: &lw.input_norm,
10198                post_norm: &lw.post_norm,
10199                ffn: crate::gpu::GraphFfn::AttentionOnly,
10200                attn: crate::gpu::GraphAttn::Full {
10201                    wq: q,
10202                    wk: k,
10203                    wv: v,
10204                    wo: o,
10205                    q_norm: q_norm.as_deref(),
10206                    k_norm: k_norm.as_deref(),
10207                    late_qk_norm: self.qk_norm_after_rope,
10208                    bias: bias
10209                        .as_ref()
10210                        .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice())),
10211                    output_gate: false,
10212                    cpu_k: self.kv_cache.layers[li].k_heads(),
10213                    cpu_v: self.kv_cache.layers[li].v_heads(),
10214                    geom: self.graph_attn_geom(li),
10215                },
10216            };
10217            let (nkv, hd, rd) = self.layer_geom(li);
10218            crate::gpu::forward_batch_graph_at(
10219                &model,
10220                self.graph_kv_id,
10221                li,
10222                &[layer],
10223                &self.inv_freq,
10224                h,
10225                self.layer_num_heads(li),
10226                nkv,
10227                hd,
10228                rd,
10229                self.hidden_size,
10230                1,
10231                positions,
10232                self.kv_cache.max_seq_len,
10233                self.norm_style == cortiq_core::NormStyle::Gemma,
10234                self.rms_eps as f32,
10235                self.attn_scale,
10236                b,
10237                &[],
10238                self.o1_epoch,
10239                None,
10240                None,
10241            )
10242        };
10243        match outcome {
10244            Out::Completed => {}
10245            Out::Declined => return Out::Declined,
10246            Out::Failed => {
10247                self.graph_failed
10248                    .store(true, std::sync::atomic::Ordering::Relaxed);
10249                return Out::Failed;
10250            }
10251        }
10252        let attn_ns = attn_started.elapsed().as_nanos() as u64;
10253        let hs = self.hidden_size;
10254        let lw = &self.weights.layers[li];
10255        let FfnKind::Moe(m) = &lw.ffn else {
10256            unreachable!()
10257        };
10258        let mut post = vec![0.0; h.len()];
10259        for (x, y) in h.chunks_exact(hs).zip(post.chunks_exact_mut(hs)) {
10260            inference::rms_norm_into(x, &lw.post_norm, self.rms_eps, self.norm_style, y);
10261        }
10262        let mut ffn = if b == 1 {
10263            moe_ffn_banked(&mut self.mimo_moe, li, m, &post, self.pool.as_deref())
10264        } else {
10265            moe_ffn_banked_rows(
10266                &mut self.mimo_moe,
10267                li,
10268                m,
10269                &post,
10270                b,
10271                hs,
10272                self.pool.as_deref(),
10273            )
10274        };
10275        for (x, &f) in h.iter_mut().zip(&ffn) {
10276            *x += f;
10277        }
10278        attention::recycle_buf(&mut ffn);
10279        if self.layer_dump.is_some() {
10280            for (&pos, row) in positions.iter().zip(h.chunks_exact(hs)) {
10281                self.dump_layer_row(pos, li, row);
10282            }
10283        }
10284        crate::mimo_moe::note_attention_graph(b, attn_ns);
10285        Out::Completed
10286    }
10287
10288    fn layer_attn_plain(&self, li: usize) -> bool {
10289        self.kv_heads_per_layer.is_none()
10290            && self.v_head_dim.is_none()
10291            && self.global_attn.is_none()
10292            && self.layer_window(li).is_none()
10293            && self.kv_cache.layers[li].sinks.is_none()
10294    }
10295
10296    /// Forward one position through all layers (hybrid dispatch).
10297    fn forward_layers(
10298        &mut self,
10299        hidden: &[f32],
10300        position: usize,
10301        task_mask: Option<&TaskMask>,
10302    ) -> Vec<f32> {
10303        let out = self.forward_layers_upto(hidden, position, task_mask, None);
10304        self.o1_progress();
10305        out
10306    }
10307
10308    // ── Network pipeline-split building blocks (coordinator/worker) ──
10309    // A remote worker owns layers [from ..= upto] and their KV; the
10310    // coordinator owns the rest plus embed / final norm / head. Attention
10311    // causality is per-layer, so a whole prompt's boundary hiddens ship
10312    // as one batch and decode ships one vector per token.
10313
10314    /// Embed one token id (embed multiplier applied).
10315    pub fn embed_id(&self, id: u32) -> Vec<f32> {
10316        self.embed_single(id)
10317    }
10318
10319    /// Refuse the archs/modes whose forward cannot be cut at a layer
10320    /// boundary. Loud by design: a split that silently changed the math
10321    /// would be a chimera.
10322    pub fn split_supported(&self) -> Result<(), String> {
10323        if self.dsv4.is_some() {
10324            return Err(
10325                "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
10326            );
10327        }
10328        if self.dsv41.is_some() {
10329            return Err(
10330                "network split: DeepSeek-V4.1 owns the shared CED/CSA2 state (not splittable)"
10331                    .into(),
10332            );
10333        }
10334        if self.qwen4_exp.is_some() {
10335            return Err(
10336                "network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
10337            );
10338        }
10339        if self.g3n.is_some() {
10340            return Err(
10341                "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
10342            );
10343        }
10344        Ok(())
10345    }
10346
10347    /// Forward `hidden` through layers [from ..= upto] at `position`,
10348    /// appending those layers' KV/state. Both split sides call this
10349    /// over their own range; a task mask applies to the span's own
10350    /// layers (each side masks what it runs).
10351    pub fn forward_span(
10352        &mut self,
10353        hidden: &[f32],
10354        position: usize,
10355        from: usize,
10356        upto: usize,
10357        task_mask: Option<&TaskMask>,
10358    ) -> Result<Vec<f32>, String> {
10359        self.split_supported()?;
10360        if from > upto || upto >= self.num_layers {
10361            return Err(format!(
10362                "forward_span: layer range {from}..={upto} outside 0..{}",
10363                self.num_layers
10364            ));
10365        }
10366        if hidden.len() != self.hidden_size {
10367            return Err(format!(
10368                "forward_span: hidden len {} ≠ hidden_size {}",
10369                hidden.len(),
10370                self.hidden_size
10371            ));
10372        }
10373        let out = self.forward_layers_span(hidden, position, task_mask, from, Some(upto));
10374        self.o1_progress();
10375        if self
10376            .graph_failed
10377            .swap(false, std::sync::atomic::Ordering::Relaxed)
10378        {
10379            self.cancel
10380                .store(false, std::sync::atomic::Ordering::Relaxed);
10381            self.clear_sequence_state();
10382            return Err("forward_span: deferred O(1) transition failed".into());
10383        }
10384        Ok(out)
10385    }
10386
10387    /// Final norm + lm_head over a boundary hidden (the final-logit
10388    /// softcap is applied by lm_head_forward itself).
10389    pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
10390        let normed = inference::rms_norm(
10391            hidden,
10392            &self.weights.final_norm,
10393            self.rms_eps,
10394            self.norm_style,
10395        );
10396        self.lm_head_forward(&normed)
10397    }
10398
10399    /// Sample the next token with this pipeline's sampler state.
10400    pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
10401        sampler::sample_with_scratch(
10402            logits,
10403            &self.sampler_config,
10404            past_tokens,
10405            &mut self.rng,
10406            &mut self.sampler_scratch,
10407        )
10408    }
10409
10410    /// Fresh sequence: clear KV, reuse history and device mirrors.
10411    pub fn reset_session(&mut self) {
10412        self.clear_sequence_state();
10413    }
10414
10415    /// Batched span prefill from token ids (coordinator side): embed +
10416    /// layers [0 ..= upto]; returns the boundary hiddens of ALL positions
10417    /// (ids.len() × hidden). Rides the same layer-major machinery as the
10418    /// local prefill; falls back to the per-position walk under
10419    /// CMF_PREFILL=seq.
10420    pub fn prefill_span_ids(
10421        &mut self,
10422        ids: &[u32],
10423        start_pos: usize,
10424        upto: usize,
10425        task_mask: Option<&TaskMask>,
10426    ) -> Result<Vec<f32>, String> {
10427        self.split_supported()?;
10428        if upto >= self.num_layers {
10429            return Err(format!(
10430                "prefill_span_ids: upto {upto} outside 0..{}",
10431                self.num_layers
10432            ));
10433        }
10434        // Same predicate as the whole-stack prefill: a span whose GDN
10435        // state lives on the device must walk positions through the
10436        // graph, not through the batched CPU span.
10437        if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10438            let out =
10439                self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1);
10440            self.check_o1_progress_failure("prefill_span_ids")?;
10441            Ok(out)
10442        } else {
10443            let hs = self.hidden_size;
10444            let mut out = Vec::with_capacity(ids.len() * hs);
10445            for (i, &id) in ids.iter().enumerate() {
10446                let emb = self.embed_id(id);
10447                out.extend_from_slice(&self.forward_span(
10448                    &emb,
10449                    start_pos + i,
10450                    0,
10451                    upto,
10452                    task_mask,
10453                )?);
10454            }
10455            Ok(out)
10456        }
10457    }
10458
10459    /// Batched span prefill from boundary hiddens (worker side): layers
10460    /// [from ..= upto] for every position in the batch; returns the batch.
10461    pub fn prefill_span_hidden(
10462        &mut self,
10463        hidden: &[f32],
10464        start_pos: usize,
10465        from: usize,
10466        upto: usize,
10467        task_mask: Option<&TaskMask>,
10468    ) -> Result<Vec<f32>, String> {
10469        self.split_supported()?;
10470        let hs = self.hidden_size;
10471        if hidden.is_empty() || hidden.len() % hs != 0 {
10472            return Err(format!(
10473                "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
10474                hidden.len()
10475            ));
10476        }
10477        if from > upto || upto >= self.num_layers {
10478            return Err(format!(
10479                "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
10480                self.num_layers
10481            ));
10482        }
10483        if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10484            let out = self.prefill_batch_span(
10485                PrefillIn::Hidden(hidden),
10486                start_pos,
10487                task_mask,
10488                from,
10489                upto + 1,
10490            );
10491            self.check_o1_progress_failure("prefill_span_hidden")?;
10492            Ok(out)
10493        } else {
10494            let b = hidden.len() / hs;
10495            let mut out = Vec::with_capacity(hidden.len());
10496            for i in 0..b {
10497                let h = self.forward_span(
10498                    &hidden[i * hs..(i + 1) * hs],
10499                    start_pos + i,
10500                    from,
10501                    upto,
10502                    task_mask,
10503                )?;
10504                out.extend_from_slice(&h);
10505            }
10506            Ok(out)
10507        }
10508    }
10509
10510    /// Build the whole-token wgpu graph for a pure-attention q1 model (every
10511    /// layer Full q1 + dense q1 FFN, no gate/bias). Returns the post-stack
10512    /// hidden (caller does final norm + lm_head), or None to fall back.
10513    fn try_token_graph_wgpu(
10514        &self,
10515        hidden: &[f32],
10516        position: usize,
10517        logits_out: &mut Vec<f32>,
10518        layers_run: &mut usize,
10519    ) -> Option<Result<Vec<f32>, ()>> {
10520        self.try_token_graph_wgpu_steps(
10521            hidden,
10522            position,
10523            logits_out,
10524            1,
10525            None,
10526            Some(layers_run),
10527            0,
10528            self.num_layers,
10529        )
10530    }
10531
10532    /// The span twin (network split): the graph covers [from..upto_excl)
10533    /// — one submit per SEGMENT per token. lm_head folds in only when
10534    /// the span reaches the last layer.
10535    fn try_token_graph_wgpu_span(
10536        &self,
10537        hidden: &[f32],
10538        position: usize,
10539        logits_out: &mut Vec<f32>,
10540        from: usize,
10541        upto_excl: usize,
10542        layers_run: &mut usize,
10543    ) -> Option<Result<Vec<f32>, ()>> {
10544        self.try_token_graph_wgpu_steps(
10545            hidden,
10546            position,
10547            logits_out,
10548            1,
10549            None,
10550            Some(layers_run),
10551            from,
10552            upto_excl,
10553        )
10554    }
10555
10556    /// Greedy burst: forward `t_next` and let the device pick + re-embed
10557    /// the next k−1 tokens — k frames, ONE submit, k ids back. The ZML
10558    /// trade, on wgpu. None ⇒ caller keeps the per-token path.
10559    fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
10560        if self.o1_active() || self.attn_softcap > 0.0 {
10561            return None;
10562        }
10563        // The burst builds the whole-token graph; attention the graph's
10564        // per-layer geometry cannot express keeps the per-token path.
10565        if let Some(reason) = self.wgpu_graph_attn_decline() {
10566            self.note_graph_decline("wgpu multi-burst", reason);
10567            return None;
10568        }
10569        let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
10570        if !graph_on || self.graph_refused() {
10571            // Same memo as the decode site: this path builds the very
10572            // same graph, so a model it cannot build for must not be
10573            // walked again here either. Missing this guard was worth
10574            // 2.5x on an Adreno — 0.361 tok/s against 0.905 — because
10575            // the burst retried per token what decode had already given
10576            // up on.
10577            return None;
10578        }
10579        let emb = self.embed_single(t_next);
10580        let mut lg = Vec::new();
10581        let mut ids = Vec::new();
10582        match self.try_token_graph_wgpu_steps(
10583            &emb,
10584            position,
10585            &mut lg,
10586            k,
10587            Some(&mut ids),
10588            None,
10589            0,
10590            self.num_layers,
10591        ) {
10592            Some(Ok(_)) => {}
10593            Some(Err(())) => {
10594                // Preserve the backend's post-admission failure through the
10595                // Option-based burst API.  The decode caller consumes this
10596                // flag and clears the sequence instead of falling through
10597                // to a stale CPU recurrent state.
10598                self.graph_failed
10599                    .store(true, std::sync::atomic::Ordering::Relaxed);
10600                return None;
10601            }
10602            None => return None,
10603        }
10604        (ids.len() == k).then_some(ids)
10605    }
10606
10607    /// Multi-step greedy: k whole frames in ONE submit, argmax and re-embed
10608    /// on the device. `ids_out` receives the k winner ids; the hidden/logits
10609    /// outputs are NOT produced in that mode.
10610    fn try_token_graph_wgpu_steps(
10611        &self,
10612        hidden: &[f32],
10613        position: usize,
10614        logits_out: &mut Vec<f32>,
10615        steps: usize,
10616        ids_out: Option<&mut Vec<u32>>,
10617        layers_run: Option<&mut usize>,
10618        from: usize,
10619        upto_excl: usize,
10620    ) -> Option<Result<Vec<f32>, ()>> {
10621        // The bank has already reserved its VRAM. Never build a second
10622        // all-expert arena across bank-owned layers (including bursts).
10623        let upto_excl = match self.mimo_moe.graph_prefix_end() {
10624            Some(end) if end < upto_excl => {
10625                if steps != 1 || layers_run.is_none() || from >= end {
10626                    return None;
10627                }
10628                end
10629            }
10630            _ => upto_excl,
10631        };
10632        // O(1) Nyström decode runs off the sealed state, not the KV cache the
10633        // graph mirrors — never take the graph while o1 is active.
10634        let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
10635        if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
10636            // Softcapped scores have no graph kernel yet — CPU owns them.
10637            // o1 rides the graph only behind CMF_O1_GPU=1 while the port
10638            // proves itself; without it the CPU path owns o1 as before.
10639            return None;
10640        }
10641        // Per-layer KV heads, narrow V, sinks and sliding windows ride
10642        // `GraphAttn::Full::geom` (the ATTEND_X kernels). Anything that
10643        // geometry cannot express declines here, by name — before the
10644        // per-layer gate existed a sliding/sink model ran the graph as
10645        // full-context attention, fluent and wrong. The caller memoizes
10646        // the refusal.
10647        if let Some(reason) = self.wgpu_graph_attn_decline() {
10648            self.note_graph_decline("wgpu token graph", reason);
10649            return None;
10650        }
10651        // Per-layer sealed o1 state for the graph. During prefill the
10652        // state is still Collecting -> views are None -> the graph
10653        // refuses below and the CPU prefill records the q trace and
10654        // seals, exactly as the o1 design requires.
10655        let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
10656            .map(|li| {
10657                if !o1_gpu {
10658                    return None;
10659                }
10660                self.kv_cache.layers[self.phys_layer(li)].o1_views()
10661            })
10662            .collect();
10663        if self.o1_active() && o1_gpu {
10664            // Any o1 layer not sealed (or degenerate exact-only) keeps the
10665            // whole token on the CPU: half-graph forwards would desync.
10666            let want: usize = (from..upto_excl)
10667                .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
10668                .count();
10669            let have = o1_views.iter().filter(|v| v.is_some()).count();
10670            if want == 0 || have != want {
10671                // The silent twin of the gpu-side o1 gates, found the
10672                // same way: a 15x decode drop with an empty log. Views
10673                // stay None until the layer's state SEALS, so `have`
10674                // lagging `want` early in a run is the o1 design working
10675                // — but it must say so, or the next reader spends a
10676                // night proving the kernels innocent.
10677                // On CHANGE, not once: the first decline is the legal
10678                // unsealed prefill, and a once-print buries the state
10679                // that matters — what the count reads AFTER the seal.
10680                use std::sync::atomic::{AtomicUsize, Ordering};
10681                static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
10682                let code = have * 1000 + want;
10683                if LAST.swap(code, Ordering::Relaxed) != code {
10684                    tracing::warn!(
10685                        "o1 graph: {have} of {want} layers sealed — per-op until all seal"
10686                    );
10687                }
10688                return None;
10689            }
10690        }
10691        let nh = self.num_heads;
10692        let (nkv, hd, rd) = self.layer_geom(0);
10693        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
10694        let mut layers = Vec::with_capacity(upto_excl - from);
10695        let mut model = None;
10696        let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
10697        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
10698            if let Some((m, i, kind, rs)) = t
10699                .graph_weight()
10700                .or_else(|| t.graph_weight_descriptor())
10701            {
10702                let name = &m.tensors[i].name;
10703                let prism = if crate::prism::is_inverse_embedding(m, name) {
10704                    crate::gpu::GraphPrismOp::InverseEmbedding
10705                } else if crate::prism::is_forward_weight(m, name) {
10706                    crate::gpu::GraphPrismOp::Forward
10707                } else {
10708                    crate::gpu::GraphPrismOp::None
10709                };
10710                return Some(crate::gpu::GraphW {
10711                    idx: i,
10712                    kind,
10713                    row_scale: rs,
10714                    data: &[],
10715                    prism,
10716                    affine: crate::prism::is_affine_target(m, name),
10717                });
10718            }
10719            // Small unquantized projections (GDN in_proj_a/b) stay f32.
10720            match t.as_f32() {
10721                Some(d) => Some(crate::gpu::GraphW {
10722                    idx: 0,
10723                    kind: 4,
10724                    row_scale: &[],
10725                    data: d,
10726                    prism: crate::gpu::GraphPrismOp::None,
10727                    affine: false,
10728                }),
10729                None => {
10730                    if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
10731                        eprintln!("batch graph: weight has no graph/f32 representation");
10732                    }
10733                    None
10734                }
10735            }
10736        }
10737        for li in from..upto_excl {
10738            let lw = &self.weights.layers[self.phys_layer(li)];
10739            if dbg {
10740                let ak = match &lw.attn {
10741                    AttnKind::Mla(_) => "Mla".into(),
10742                    AttnKind::Full {
10743                        output_gate, bias, ..
10744                    } => format!("Full gate={output_gate} bias={}", bias.is_some()),
10745                    AttnKind::LinearGdn(_) => "LinearGdn".into(),
10746                    AttnKind::Kda(_) => "Kda".into(),
10747                    AttnKind::Linear(_) => "Linear".into(),
10748                    AttnKind::ShortConv(_) => "ShortConv".into(),
10749                    AttnKind::Bounded(_) => "Bounded".into(),
10750                };
10751                let fk = match &lw.ffn {
10752                    FfnKind::Dense(_) => "Dense",
10753                    FfnKind::Moe(_) => "Moe",
10754                    FfnKind::DenseMoe(_) => "DenseMoe",
10755                };
10756                eprintln!("graph L{li}: attn={ak} ffn={fk}");
10757            }
10758            let gffn = match &lw.ffn {
10759                FfnKind::DenseMoe(_) => return None, // dual branch: CPU path
10760                // A tube layer is several matrices, not one — the
10761                // whole-layer graph has no shape for it yet.
10762                FfnKind::Dense(d) if !d.segs.is_empty() => return None,
10763                FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
10764                    gate: gw(&d.gate_proj)?,
10765                    up: gw(&d.up_proj)?,
10766                    down: gw(&d.down_proj)?,
10767                },
10768                FfnKind::Moe(m) => {
10769                    // Adaptive τ and expert masks keep the CPU path, where
10770                    // they are implemented. Sigmoid routing with a selection
10771                    // bias (LFM2-MoE / DeepSeek noaux_tc), a routed scale ≠ 1
10772                    // and an UNGATED shared expert (HunYuan hy_v3: ×2.826 on
10773                    // the routed mix, the shared expert at weight 1) are all
10774                    // graphed — before, every such token fell to the per-op
10775                    // path whole (145 submits/token on Hy-MT2-30B-A3B).
10776                    if m.route_tau.is_some() || m.mask.is_some() {
10777                        return None;
10778                    }
10779                    let shared = m.shared.as_ref();
10780                    let has_shared = shared.is_some();
10781                    let shared_gated = matches!(shared, Some((_, Some(_))));
10782                    let sgate = match shared {
10783                        Some((_, Some(sg))) => gw(sg)?,
10784                        // No gate (hy_v3) or no shared expert at all: the
10785                        // router weight stands in so the plumbing stays
10786                        // total; the select kernels pin weight 1 or skip.
10787                        _ => gw(&m.router)?,
10788                    };
10789                    let router = gw(&m.router)?;
10790                    // The resident MoE kernels do not yet carry the
10791                    // descriptor-aware transform through router/shared-gate
10792                    // selection.  Refuse the complete layer instead of
10793                    // scoring with an untransformed Prism plane (the dense
10794                    // path has an explicit FWHT boundary below).
10795                    if router.prism != crate::gpu::GraphPrismOp::None
10796                        || sgate.prism != crate::gpu::GraphPrismOp::None
10797                        || router.affine
10798                        || sgate.affine
10799                    {
10800                        tracing::warn!(
10801                            "resident MoE declined: Prism/affine router or shared gate transform is not implemented"
10802                        );
10803                        return None;
10804                    }
10805                    let inter = m.experts.first()?.gate_proj.rows();
10806                    let mut experts = Vec::with_capacity(m.experts.len() + 1);
10807                    // q4t or q4tp, but not both in one layer — the kernels
10808                    // are picked per layer, not per expert.
10809                    let mut q4tp: Option<bool> = None;
10810                    // The mixed 2-bit profile: q2tp gate/up over a q4tp
10811                    // down. Uniform across the layer, like `q4tp` itself.
10812                    let mut gu_q2: Option<bool> = None;
10813                    for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
10814                        if !matches!(e.act, Act::Silu)
10815                            || e.gate_proj.rows() != inter
10816                            || e.up_proj.rows() != inter
10817                        {
10818                            return None;
10819                        }
10820                        // Expert tensors are packed into one resident buffer
10821                        // and the MoE kernels have no transform slot per
10822                        // expert.  Keep the CPU/per-op owner for Prism or
10823                        // affine experts rather than silently using raw bytes.
10824                        for expert_weight in [&e.gate_proj, &e.up_proj, &e.down_proj] {
10825                            let Some((em, ei, _, _)) = expert_weight
10826                                .graph_weight()
10827                                .or_else(|| expert_weight.graph_weight_descriptor())
10828                            else {
10829                                return None;
10830                            };
10831                            let name = &em.tensors[ei].name;
10832                            if crate::prism::is_forward_weight(em, name)
10833                                || crate::prism::is_inverse_embedding(em, name)
10834                                || crate::prism::is_affine_target(em, name)
10835                            {
10836                                tracing::warn!(
10837                                    "resident MoE declined: expert Prism/affine transform is not implemented"
10838                                );
10839                                return None;
10840                            }
10841                        }
10842                        let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
10843                            Some((mm, gi)) => (
10844                                mm,
10845                                gi,
10846                                e.up_proj.mapped_q4t()?.1,
10847                                e.down_proj.mapped_q4t()?.1,
10848                                false,
10849                                false,
10850                            ),
10851                            None => match e.gate_proj.mapped_q2tp() {
10852                                Some((mm, gi)) => (
10853                                    mm,
10854                                    gi,
10855                                    e.up_proj.mapped_q2tp()?.1,
10856                                    e.down_proj.mapped_q4tp()?.1,
10857                                    true,
10858                                    true,
10859                                ),
10860                                None => {
10861                                    let (mm, gi) = e.gate_proj.mapped_q4tp()?;
10862                                    (
10863                                        mm,
10864                                        gi,
10865                                        e.up_proj.mapped_q4tp()?.1,
10866                                        e.down_proj.mapped_q4tp()?.1,
10867                                        true,
10868                                        false,
10869                                    )
10870                                }
10871                            },
10872                        };
10873                        if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
10874                        {
10875                            // The shared expert rides in the same packed
10876                            // buffer as the routed ones, so a layer that
10877                            // mixes layouts cannot be indexed by one stride.
10878                            // Say so: the symptom is a whole model quietly
10879                            // running its MoE on the CPU.
10880                            tracing::warn!(
10881                                "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."
10882                            );
10883                            return None;
10884                        }
10885                        model.get_or_insert_with(|| mm.clone());
10886                        experts.push((gi, ui, di));
10887                    }
10888                    crate::gpu::GraphFfn::Moe {
10889                        router,
10890                        shared_gate: sgate,
10891                        experts,
10892                        n_exp: m.experts.len(),
10893                        // CMF_TOPK_PROBE: timing probe only — output is WRONG.
10894                        // Fewer experts shrink the MoE arithmetic while the
10895                        // dispatch count stays identical, which is the only
10896                        // clean way to tell a launch-bound decode from a
10897                        // compute-bound one.
10898                        top_k: std::env::var("CMF_TOPK_PROBE")
10899                            .ok()
10900                            .and_then(|v| v.parse::<usize>().ok())
10901                            .filter(|k| *k > 0 && *k <= m.top_k)
10902                            .unwrap_or(m.top_k),
10903                        inter,
10904                        norm_topk: m.norm_topk_prob,
10905                        q4tp: q4tp?,
10906                        gu_q2: gu_q2.unwrap_or(false),
10907                        sigmoid: m.router_sigmoid,
10908                        bias: m.expert_bias.as_deref(),
10909                        has_shared,
10910                        shared_gated,
10911                        route_scale: m.routed_scaling,
10912                    }
10913                }
10914            };
10915            let attn = match &lw.attn {
10916                AttnKind::Full {
10917                    wq,
10918                    wk,
10919                    wv,
10920                    wo,
10921                    q_norm,
10922                    k_norm,
10923                    output_gate,
10924                    softplus_gate,
10925                    bias,
10926                } => {
10927                    if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
10928                        return None;
10929                    }
10930                    let (m, _, _, _) = wq
10931                        .graph_weight()
10932                        .or_else(|| wq.graph_weight_descriptor())?;
10933                    model = Some(m.clone());
10934                    crate::gpu::GraphAttn::Full {
10935                        wq: gw(wq)?,
10936                        wk: gw(wk)?,
10937                        wv: gw(wv)?,
10938                        wo: gw(wo)?,
10939                        q_norm: q_norm.as_deref(),
10940                        k_norm: k_norm.as_deref(),
10941                        late_qk_norm: self.qk_norm_after_rope,
10942                        bias: bias
10943                            .as_ref()
10944                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
10945                        output_gate: *output_gate,
10946                        cpu_k: self.kv_cache.layers[li].k_heads(),
10947                        cpu_v: self.kv_cache.layers[li].v_heads(),
10948                        geom: self.graph_attn_geom(li),
10949                    }
10950                }
10951                AttnKind::LinearGdn(w) => {
10952                    let cfg = self.gdn_cfg?;
10953                    let (m, _, _, _) = w
10954                        .in_proj_qkv
10955                        .graph_weight()
10956                        .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
10957                    model = Some(m.clone());
10958                    crate::gpu::GraphAttn::Gdn {
10959                        qkv: gw(&w.in_proj_qkv)?,
10960                        z: gw(&w.in_proj_z)?,
10961                        a: gw(&w.in_proj_a)?,
10962                        b: gw(&w.in_proj_b)?,
10963                        out: gw(&w.out_proj)?,
10964                        conv1d: &w.conv1d,
10965                        a_log: &w.a_log,
10966                        dt_bias: &w.dt_bias,
10967                        norm: &w.norm,
10968                        nv: cfg.num_v_heads,
10969                        nk: cfg.num_k_heads,
10970                        dk: cfg.key_head_dim,
10971                        dv: cfg.value_head_dim,
10972                        kk: cfg.conv_kernel,
10973                        cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
10974                    }
10975                }
10976                AttnKind::ShortConv(w) => {
10977                    let cfg = self.short_conv_cfg?;
10978                    let (m, _, _, _) = w
10979                        .in_proj
10980                        .graph_weight()
10981                        .or_else(|| w.in_proj.graph_weight_descriptor())?;
10982                    model = Some(m.clone());
10983                    crate::gpu::GraphAttn::ShortConv {
10984                        inp: gw(&w.in_proj)?,
10985                        out: gw(&w.out_proj)?,
10986                        taps: &w.conv,
10987                        kernel: cfg.kernel,
10988                        cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
10989                    }
10990                }
10991                _ => return None,
10992            };
10993            layers.push(crate::gpu::GraphLayer {
10994                input_norm: &lw.input_norm,
10995                attn,
10996                post_norm: &lw.post_norm,
10997                ffn: gffn,
10998            });
10999        }
11000        let model = model?;
11001        // Fold final-norm + lm_head into the graph when this call wants logits
11002        // and the lm_head is a graphable (quantized) weight — the graph then
11003        // reads back logits (into logits_out) instead of the hidden, dropping
11004        // the separate CPU/GPU lm_head op + its sync. Never the f32 fallback:
11005        // an unquantized lm_head is vocab·hidden and must not be uploaded.
11006        let lm_gw = if upto_excl == self.num_layers
11007            && self.graph_want_logits
11008            && std::env::var("CMF_GPU_LMHEAD")
11009                .map(|v| v != "0")
11010                .unwrap_or(true)
11011        {
11012            self.weights
11013                .lm_head
11014                .graph_weight()
11015                .or_else(|| self.weights.lm_head.graph_weight_descriptor())
11016                .map(|(m, i, kind, rs)| {
11017                let name = &m.tensors[i].name;
11018                let prism = if crate::prism::is_inverse_embedding(m, name) {
11019                    crate::gpu::GraphPrismOp::InverseEmbedding
11020                } else if crate::prism::is_forward_weight(m, name) {
11021                    crate::gpu::GraphPrismOp::Forward
11022                } else {
11023                    crate::gpu::GraphPrismOp::None
11024                };
11025                (
11026                    crate::gpu::GraphW {
11027                        idx: i,
11028                        kind,
11029                        row_scale: rs,
11030                        data: &[],
11031                        prism,
11032                        affine: crate::prism::is_affine_target(m, name),
11033                    },
11034                    self.weights.lm_head.rows(),
11035                )
11036            })
11037        } else {
11038            None
11039        };
11040        let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
11041        // Multi-step re-embeds the winner on the device.
11042        let emb_gw = if steps > 1 {
11043            self.weights
11044                .embed_tokens
11045                .graph_weight()
11046                .or_else(|| self.weights.embed_tokens.graph_weight_descriptor())
11047                .map(|(m, i, kind, rs)| {
11048                    let name = &m.tensors[i].name;
11049                    let prism = if crate::prism::is_inverse_embedding(m, name) {
11050                        crate::gpu::GraphPrismOp::InverseEmbedding
11051                    } else if crate::prism::is_forward_weight(m, name) {
11052                        crate::gpu::GraphPrismOp::Forward
11053                    } else {
11054                        crate::gpu::GraphPrismOp::None
11055                    };
11056                    (
11057                        crate::gpu::GraphW {
11058                            idx: i,
11059                            kind,
11060                            row_scale: rs,
11061                            data: &[],
11062                            prism,
11063                            affine: crate::prism::is_affine_target(m, name),
11064                        },
11065                        self.weights.embed_tokens.rows(),
11066                        self.embed_multiplier,
11067                    )
11068                })
11069        } else {
11070            None
11071        };
11072
11073        // Loop boundaries: virtual layer indices after which final_norm is
11074        // applied (mid-stack only; the GLOBAL last layer's norm folds into
11075        // lm_head). Span-relative — the executor compares its enumerate
11076        // index. A span ending mid-stack keeps its boundary norm even when
11077        // it is the span's own last layer.
11078        let loop_norm_at: Vec<usize> = if self.loop_final_norm {
11079            (from..upto_excl.min(self.num_layers - 1))
11080                .filter(|&li| (li + 1) % self.physical_layers == 0)
11081                .map(|li| li - from)
11082                .collect()
11083        } else {
11084            Vec::new()
11085        };
11086        let mut h = hidden.to_vec();
11087        // The normal decode path only needs the fused lm-head logits.  A
11088        // CMF_LOGIT_DUMP diagnostic, however, promises a prompt-boundary
11089        // post-stack hidden alongside those logits; request the existing
11090        // second readback only for that explicit probe instead of dumping
11091        // the input copy left in `h` by a folded-head graph.
11092        let dump_hidden = std::env::var_os("CMF_LOGIT_DUMP").is_some();
11093        let outcome = crate::gpu::forward_token_graph(
11094            &model,
11095            self.graph_kv_id,
11096            &layers,
11097            &o1_views,
11098            self.o1_epoch,
11099            &self.inv_freq,
11100            &mut h,
11101            nh,
11102            nkv,
11103            hd,
11104            self.attn_scale,
11105            rd,
11106            self.hidden_size,
11107            self.intermediate_size,
11108            position,
11109            self.kv_cache.max_seq_len,
11110            gemma,
11111            self.rms_eps as f32,
11112            lm,
11113            &self.weights.final_norm,
11114            logits_out,
11115            &loop_norm_at,
11116            steps,
11117            emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
11118            ids_out,
11119            layers_run,
11120            from,
11121            dump_hidden,
11122        );
11123        match outcome {
11124            crate::gpu::TokenGraphOutcome::Completed => Some(Ok(h)),
11125            crate::gpu::TokenGraphOutcome::Failed => Some(Err(())),
11126            crate::gpu::TokenGraphOutcome::Declined => None,
11127        }
11128    }
11129
11130    /// Batched prefill: k contiguous prompt positions through the whole wgpu
11131    /// graph in ONE submit (projections/FFN as GEMMs). `hiddens` is [k·hidden]
11132    /// in/out (embeddings in, layer output out); KV mirror / GDN state advance.
11133    /// false ⇒ unsupported → caller keeps the per-position graph.
11134    /// The b-row Metal graph plan for the whole model: every layer as a
11135    /// GDN run or a full-attention item, all-or-nothing (a layer outside the
11136    /// graph's contract → None, the caller runs plain). Shared by the
11137    /// speculative verify and the batched prefill.
11138    #[cfg(target_os = "macos")]
11139    #[allow(clippy::type_complexity)]
11140    fn metal_rows_plan(
11141        &self,
11142    ) -> Option<(
11143        Vec<MetalRowsItem<'_>>,
11144        std::sync::Arc<cortiq_core::CmfModel>,
11145        Option<crate::gpu_metal::GdnGpuCfg>,
11146    )> {
11147        use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
11148        let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
11149        if !graph_force
11150            || !crate::gpu::enabled_here()
11151            || std::env::var("CMF_GPU_BLOCK")
11152                .map(|v| v == "0")
11153                .unwrap_or(false)
11154            || self.attn_softcap > 0.0
11155            || self.o1_active()
11156            || self.swa.is_some()
11157            || self.global_attn.is_some()
11158            || self.attention_heads_per_layer.is_some()
11159            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
11160            || self.graph_attn_decline_reason().is_some()
11161            || self.attn_v_norm
11162            || self.loop_final_norm
11163        {
11164            return None;
11165        }
11166        let attend_contract = self.head_dim % 4 == 0
11167            && self.head_dim <= 256
11168            && self.rotary_dim >= 2
11169            && self.rotary_dim <= self.head_dim
11170            && (self.rotary_dim / 2) % 32 == 0
11171            && self.num_kv_heads > 0
11172            && self.num_heads % self.num_kv_heads == 0;
11173        if !attend_contract {
11174            return None;
11175        }
11176        let mut plan: Vec<MetalRowsItem> = Vec::new();
11177        let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
11178        for li in 0..self.num_layers {
11179            let lw = &self.weights.layers[self.phys_layer(li)];
11180            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
11181                return None;
11182            }
11183            let ffn = match &lw.ffn {
11184                FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
11185                    let (Some(g), Some(u), Some(dn)) = (
11186                        d.gate_proj.metal_graph_parts(),
11187                        d.up_proj.metal_graph_parts(),
11188                        d.down_proj.metal_graph_parts(),
11189                    ) else {
11190                        return None;
11191                    };
11192                    MetalFfn::Dense {
11193                        gate: g,
11194                        up: u,
11195                        down: dn,
11196                    }
11197                }
11198                _ => return None,
11199            };
11200            match &lw.attn {
11201                AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
11202                    let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
11203                        w.in_proj_qkv.metal_graph_parts(),
11204                        w.in_proj_z.metal_graph_parts(),
11205                        w.in_proj_a.f32_parts(),
11206                        w.in_proj_b.f32_parts(),
11207                        w.out_proj.metal_graph_parts(),
11208                    ) else {
11209                        return None;
11210                    };
11211                    if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
11212                        model_ref.get_or_insert_with(|| model.clone());
11213                    }
11214                    let gl = GdnGpuLayer {
11215                        attn_norm: &lw.input_norm,
11216                        post_norm: &lw.post_norm,
11217                        qkv,
11218                        z,
11219                        a,
11220                        b: bb,
11221                        out,
11222                        ffn,
11223                        conv1d: &w.conv1d,
11224                        a_log: &w.a_log,
11225                        dt_bias: &w.dt_bias,
11226                        gnorm: &w.norm,
11227                    };
11228                    match plan.last_mut() {
11229                        Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
11230                        _ => plan.push(MetalRowsItem::Gdn {
11231                            run: vec![gl],
11232                            first: li,
11233                        }),
11234                    }
11235                }
11236                AttnKind::Full {
11237                    wq,
11238                    wk,
11239                    wv,
11240                    wo,
11241                    q_norm,
11242                    k_norm,
11243                    output_gate,
11244                    softplus_gate: None,
11245                    bias: None,
11246                } => {
11247                    let (Some(pq), Some(pk), Some(pv), Some(po)) =
11248                        (
11249                            wq.metal_graph_parts(),
11250                            wk.metal_graph_parts(),
11251                            wv.metal_graph_parts(),
11252                            wo.metal_graph_parts(),
11253                        )
11254                    else {
11255                        return None;
11256                    };
11257                    if let QTensor::Mapped { model, .. } = wq {
11258                        model_ref.get_or_insert_with(|| model.clone());
11259                    }
11260                    let cache = &self.kv_cache.layers[li];
11261                    if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
11262                        return None;
11263                    }
11264                    plan.push(MetalRowsItem::Attn {
11265                        l: AttnGpuLayer {
11266                            attn_norm: &lw.input_norm,
11267                            post_norm: &lw.post_norm,
11268                            wq: pq,
11269                            wk: pk,
11270                            wv: pv,
11271                            wo: po,
11272                            ffn,
11273                        },
11274                        li,
11275                        q_norm: q_norm.as_deref(),
11276                        k_norm: k_norm.as_deref(),
11277                        output_gate: *output_gate,
11278                    });
11279                }
11280                _ => return None,
11281            }
11282        }
11283        let model = model_ref?;
11284        let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
11285            nv: cfg.num_v_heads,
11286            nk: cfg.num_k_heads,
11287            dk: cfg.key_head_dim,
11288            dv: cfg.value_head_dim,
11289            kk: cfg.conv_kernel,
11290            hidden: self.hidden_size,
11291            inter: self.intermediate_size,
11292            c_dim: cfg.conv_dim(),
11293            eps: cfg.rms_eps as f32,
11294            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11295        });
11296        Some((plan, model, gcfg))
11297    }
11298
11299    /// `AttnDeviceParams` for a plan item over the CPU cache as it stands.
11300    #[cfg(target_os = "macos")]
11301    #[allow(clippy::too_many_arguments)]
11302    fn metal_attn_params<'a>(
11303        li: usize,
11304        cache: &'a crate::kv_cache::LayerKvCache,
11305        q_norm: Option<&'a [f32]>,
11306        k_norm: Option<&'a [f32]>,
11307        output_gate: bool,
11308        inv_freq: &'a [f32],
11309        geom: (usize, usize, usize, usize),
11310        pos0: usize,
11311        kv_id: u64,
11312        scale: f32,
11313        eps: f32,
11314        gemma: bool,
11315        late_qk_norm: bool,
11316    ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
11317        let (nh, nkv, hd, rd) = geom;
11318        let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11319        let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11320        let cpu_stored = cpu_k[0].len() / hd;
11321        (
11322            crate::gpu_metal::AttnDeviceParams {
11323                kv_id,
11324                layer: li,
11325                nh,
11326                nkv,
11327                hd,
11328                rd,
11329                position: pos0,
11330                scale,
11331                eps,
11332                gemma,
11333                late_qk_norm,
11334                output_gate,
11335                q_norm,
11336                k_norm,
11337                inv_freq,
11338                cpu_k,
11339                cpu_v,
11340                cpu_stored,
11341                o1: None,
11342            },
11343            cpu_stored,
11344        )
11345    }
11346
11347    /// Run the rows plan over `hiddens` (b rows at `pos0..`): validate,
11348    /// encode every item, optionally the head, sync. Returns the graph
11349    /// (for the commit / state finish) plus the GDN layer indices and the
11350    /// attention layers with the row count they were encoded against.
11351    #[cfg(target_os = "macos")]
11352    #[allow(clippy::type_complexity)]
11353    fn metal_rows_run(
11354        &mut self,
11355        hiddens: &mut [f32],
11356        pos0: usize,
11357        b: usize,
11358        prefill: bool,
11359        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11360        // Greedy verify: (row length scored, the b argmax ids out) — the
11361        // head's argmax runs on the device and the logits plane is NOT
11362        // read back (`spec.2` stays empty).
11363        mut argmax_out: Option<(usize, &mut Vec<u32>)>,
11364    ) -> MetalRowsRun {
11365        use crate::gpu_metal::{GraphDims, VerifyGraph};
11366        // The previous round's commit may still be replaying into the
11367        // trunk GDN owners on the second queue: this graph reads them
11368        // (zero-copy wraps) and may reallocate them below — collect the
11369        // replay first. Normally already complete (the draft chain ran
11370        // in between); a failed replay is terminal like a failed commit.
11371        if !crate::gpu_metal::wait_replay() {
11372            tracing::error!("Metal rows graph: the pending async replay failed");
11373            return MetalRowsRun::Failed;
11374        }
11375        spec_stamp("v.wait");
11376        // Seed the GDN recurrent records the rows graph reads — GDN layers
11377        // ONLY (the record the CPU path would allocate anyway).  Sizing
11378        // every layer here planted a zero GDN-sized record on a bounded
11379        // anchor of a natively bounded file this path then refused, and
11380        // that record counted as recurrent state on macOS.
11381        let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
11382        if want > 0 {
11383            let phys = self.physical_layers.max(1);
11384            for (li, l) in self.kv_cache.layers.iter_mut().enumerate() {
11385                let is_gdn = self
11386                    .weights
11387                    .layers
11388                    .get(li % phys)
11389                    .is_some_and(|lw| matches!(lw.attn, AttnKind::LinearGdn(_)));
11390                if is_gdn && l.linear_state.len() != want {
11391                    l.linear_state = vec![0f32; want];
11392                }
11393            }
11394        }
11395        let Some((plan, model, gcfg)) = self.metal_rows_plan() else {
11396            return MetalRowsRun::Declined;
11397        };
11398        spec_stamp("v.plan");
11399        let dims = GraphDims {
11400            hidden: self.hidden_size,
11401            eps: self.rms_eps as f32,
11402            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11403        };
11404        let Some(mut graph) = (if prefill {
11405            VerifyGraph::new_prefill(&model, dims, hiddens, b)
11406        } else {
11407            VerifyGraph::new(&model, dims, hiddens, b)
11408        }) else {
11409            return MetalRowsRun::Declined;
11410        };
11411        let geom = (
11412            self.num_heads,
11413            self.num_kv_heads,
11414            self.head_dim,
11415            self.rotary_dim,
11416        );
11417        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11418        let eps = self.rms_eps as f32;
11419        let kv_id = self.graph_kv_id;
11420        let inv_freq = self.inv_freq.clone();
11421        for item in &plan {
11422            let ok = match item {
11423                MetalRowsItem::Gdn { run, .. } => gcfg
11424                    .as_ref()
11425                    .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
11426                    .unwrap_or(false),
11427                MetalRowsItem::Attn {
11428                    l,
11429                    li,
11430                    q_norm,
11431                    k_norm,
11432                    output_gate,
11433                } => {
11434                    let (p, _) = Self::metal_attn_params(
11435                        *li,
11436                        &self.kv_cache.layers[*li],
11437                        *q_norm,
11438                        *k_norm,
11439                        *output_gate,
11440                        &inv_freq,
11441                        geom,
11442                        pos0,
11443                        kv_id,
11444                        self.attn_scale,
11445                        eps,
11446                        gemma,
11447                        self.qk_norm_after_rope,
11448                    );
11449                    graph.attn_ok(l, &p)
11450                }
11451            };
11452            if !ok {
11453                use std::sync::atomic::{AtomicBool, Ordering};
11454                static SAID: AtomicBool = AtomicBool::new(false);
11455                if !SAID.swap(true, Ordering::Relaxed) {
11456                    tracing::warn!("metal rows graph: a layer failed preflight — declining");
11457                }
11458                return MetalRowsRun::Declined;
11459            }
11460        }
11461        let lm = match &spec {
11462            Some((lm, _, _)) => {
11463                if !graph.lm_head_ok(*lm) {
11464                    return MetalRowsRun::Declined;
11465                }
11466                Some(*lm)
11467            }
11468            None => None,
11469        };
11470        let mut gdn_layers = Vec::new();
11471        let mut attn_layers = Vec::new();
11472        for item in &plan {
11473            match item {
11474                MetalRowsItem::Gdn { run, first } => {
11475                    let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
11476                        .iter()
11477                        .map(|l| l.linear_state.as_slice())
11478                        .collect();
11479                    if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
11480                        return MetalRowsRun::Declined;
11481                    }
11482                    gdn_layers.extend(*first..*first + run.len());
11483                }
11484                MetalRowsItem::Attn {
11485                    l,
11486                    li,
11487                    q_norm,
11488                    k_norm,
11489                    output_gate,
11490                } => {
11491                    let (p, cpu_stored) = Self::metal_attn_params(
11492                        *li,
11493                        &self.kv_cache.layers[*li],
11494                        *q_norm,
11495                        *k_norm,
11496                        *output_gate,
11497                        &inv_freq,
11498                        geom,
11499                        pos0,
11500                        kv_id,
11501                        self.attn_scale,
11502                        eps,
11503                        gemma,
11504                        self.qk_norm_after_rope,
11505                    );
11506                    if !graph.encode_attn_b(l, &p) {
11507                        return MetalRowsRun::Declined;
11508                    }
11509                    attn_layers.push((*li, cpu_stored));
11510                }
11511            }
11512        }
11513        if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
11514            if !graph.encode_lm_head_b(final_norm, lm) {
11515                return MetalRowsRun::Declined;
11516            }
11517            // The device argmax is an OPTIMISATION, never a reason to
11518            // decline the round: if it will not encode, drop it and read
11519            // the logits plane back the old way (the head is encoded
11520            // either way, so the rows are there).
11521            if let Some((n, _)) = argmax_out.as_ref() {
11522                if !graph.encode_argmax_b(*n) {
11523                    argmax_out = None;
11524                }
11525            }
11526        }
11527        spec_stamp("v.enc");
11528        if !graph.sync() {
11529            return MetalRowsRun::Failed;
11530        }
11531        spec_stamp("v.gpu");
11532        match (spec, argmax_out) {
11533            (Some(_), Some((_, ids))) => {
11534                ids.resize(b, 0);
11535                if !graph.read_argmax(ids) {
11536                    return MetalRowsRun::Failed;
11537                }
11538                spec_stamp("v.am");
11539            }
11540            (Some((lm, _, logits)), None) => {
11541                logits.resize(b * lm.1, 0.0);
11542                if !graph.read_logits(logits) {
11543                    return MetalRowsRun::Failed;
11544                }
11545                spec_stamp("v.lg");
11546            }
11547            (None, _) => {}
11548        }
11549        if !graph.read_hidden(hiddens) {
11550            return MetalRowsRun::Failed;
11551        }
11552        spec_stamp("v.hid");
11553        MetalRowsRun::Completed(MetalVerifyPending {
11554            graph,
11555            gdn_layers,
11556            attn_layers,
11557        })
11558    }
11559
11560    /// Native-Metal twin of `try_batch_graph_wgpu`: the b rows through the
11561    /// whole model on the `VerifyGraph` (one submit), the head folded in
11562    /// when `spec` asks; `hiddens` come back as the last layer's output
11563    /// rows, `spec.2` as `[b][lm_rows]` logits. The graph is parked in
11564    /// `metal_verify` for `metal_verify_commit`.
11565    #[cfg(target_os = "macos")]
11566    fn try_batch_graph_metal(
11567        &mut self,
11568        hiddens: &mut [f32],
11569        positions: &[usize],
11570        b: usize,
11571        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11572        argmax_out: Option<(usize, &mut Vec<u32>)>,
11573    ) -> crate::gpu::BatchGraphOutcome {
11574        let _t0 = std::time::Instant::now();
11575        if positions.len() != b
11576            || positions.windows(2).any(|w| w[1] != w[0] + 1)
11577            || hiddens.len() != b * self.hidden_size
11578        {
11579            return crate::gpu::BatchGraphOutcome::Declined;
11580        }
11581        let pending = match self.metal_rows_run(hiddens, positions[0], b, false, spec, argmax_out) {
11582            MetalRowsRun::Declined => return crate::gpu::BatchGraphOutcome::Declined,
11583            MetalRowsRun::Failed => return crate::gpu::BatchGraphOutcome::Failed,
11584            MetalRowsRun::Completed(pending) => pending,
11585        };
11586        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
11587            eprintln!(
11588                "metal-verify: {:.1} ms | b={b}",
11589                _t0.elapsed().as_secs_f64() * 1e3
11590            );
11591        }
11592        self.metal_verify = Some(pending);
11593        crate::gpu::BatchGraphOutcome::Completed
11594    }
11595
11596    /// Batched prefill on the Metal rows graph: `ids` (≤ 512) at
11597    /// `start_pos..`, states written in place, K/V rows appended to the
11598    /// CPU caches; optional final norm/head logits are returned in `spec`.
11599    /// Declined means no command buffer was admitted; Failed is terminal.
11600    #[cfg(target_os = "macos")]
11601    fn prefill_rows_metal(
11602        &mut self,
11603        ids: &[u32],
11604        start_pos: usize,
11605        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11606    ) -> MetalPrefillOutcome {
11607        let b = ids.len();
11608        if b == 0 || b > 512 {
11609            return MetalPrefillOutcome::Declined;
11610        }
11611        METAL_PREFILL_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11612        let with_head = spec.is_some();
11613        let hs = self.hidden_size;
11614        let mut hiddens = vec![0f32; b * hs];
11615        for (j, &id) in ids.iter().enumerate() {
11616            let e = self.embed_single(id);
11617            hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
11618        }
11619        let mut pending = match self.metal_rows_run(&mut hiddens, start_pos, b, true, spec, None) {
11620            MetalRowsRun::Declined => return MetalPrefillOutcome::Declined,
11621            MetalRowsRun::Failed => {
11622                METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11623                return MetalPrefillOutcome::Failed;
11624            }
11625            MetalRowsRun::Completed(pending) => pending,
11626        };
11627        // states are final: copy them to the owners
11628        let idxs = pending.gdn_layers.clone();
11629        let mut outs: Vec<&mut [f32]> = self
11630            .kv_cache
11631            .layers
11632            .iter_mut()
11633            .enumerate()
11634            .filter(|(i, _)| idxs.binary_search(i).is_ok())
11635            .map(|(_, l)| l.linear_state.as_mut_slice())
11636            .collect();
11637        if !pending.graph.finish_states(&mut outs) {
11638            METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11639            return MetalPrefillOutcome::Failed;
11640        }
11641        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
11642        // Read every layer before mutating any CPU cache.  A missing mirror
11643        // row is a terminal graph failure, not a reason to append a partial
11644        // prefix and replay the remainder serially.
11645        let mut rows = Vec::with_capacity(pending.attn_layers.len());
11646        for (li, cpu_stored) in &pending.attn_layers {
11647            let mut kbuf = vec![0f32; b * nkv * hd];
11648            let mut vbuf = vec![0f32; b * nkv * hd];
11649            if !crate::gpu_metal::kv_mirror_read_rows(
11650                self.graph_kv_id,
11651                *li,
11652                nkv,
11653                hd,
11654                *cpu_stored,
11655                b,
11656                &mut kbuf,
11657                &mut vbuf,
11658            ) {
11659                METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11660                return MetalPrefillOutcome::Failed;
11661            }
11662            rows.push((*li, *cpu_stored, kbuf, vbuf));
11663        }
11664        for (li, cpu_stored, kbuf, vbuf) in rows {
11665            let cache = &mut self.kv_cache.layers[li];
11666            for r in 0..b {
11667                cache.append(
11668                    &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
11669                    &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
11670                    &[],
11671                );
11672            }
11673            crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + b);
11674        }
11675        METAL_PREFILL_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11676        if with_head {
11677            METAL_PREFILL_HEAD_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11678        }
11679        MetalPrefillOutcome::Completed(hiddens)
11680    }
11681
11682    #[cfg(target_os = "macos")]
11683    fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> MetalPrefillOutcome {
11684        self.prefill_rows_metal(ids, start_pos, None)
11685    }
11686
11687    /// Exact teacher-forced NLL through the ordinary Metal rows graph.  This
11688    /// is intentionally separate from the serial TokenGraph scorer: every
11689    /// chunk owns a real b-row graph/head completion and the recurrent/KV
11690    /// handoff is committed before the next chunk begins.
11691    #[cfg(target_os = "macos")]
11692    fn nll_batch_metal(&mut self, ids: &[u32], start: usize) -> MetalBatchNllOutcome {
11693        if ids.len() < 2 || self.o1_active() || self.head_clusters.is_some() {
11694            return MetalBatchNllOutcome::Declined;
11695        }
11696        let Some(lm) = self.weights.lm_head.metal_graph_parts() else {
11697            return MetalBatchNllOutcome::Declined;
11698        };
11699        let chunk = std::env::var("CMF_METAL_PREFILL_CHUNK")
11700            .ok()
11701            .and_then(|v| v.parse::<usize>().ok())
11702            .filter(|&v| (1..=512).contains(&v))
11703            .unwrap_or(32);
11704        let final_norm = self.weights.final_norm.clone();
11705        let mut nll = 0.0f64;
11706        let mut count = 0usize;
11707        let mut pos = 0usize;
11708        let mut completed = 0usize;
11709        while pos < ids.len() {
11710            let end = (pos + chunk).min(ids.len());
11711            let mut logits = Vec::new();
11712            let outcome = self.prefill_rows_metal(
11713                &ids[pos..end],
11714                pos,
11715                Some((lm, &final_norm, &mut logits)),
11716            );
11717            match outcome {
11718                MetalPrefillOutcome::Declined => {
11719                    return if completed == 0 {
11720                        MetalBatchNllOutcome::Declined
11721                    } else {
11722                        MetalBatchNllOutcome::Failed(format!(
11723                            "ordinary Metal NLL batch declined after {completed} chunks"
11724                        ))
11725                    };
11726                }
11727                MetalPrefillOutcome::Failed => {
11728                    return MetalBatchNllOutcome::Failed(
11729                        "ordinary Metal NLL batch failed after admission".to_string(),
11730                    );
11731                }
11732                MetalPrefillOutcome::Completed(_) => {}
11733            }
11734            completed += 1;
11735            let vocab = self.vocab_size.min(lm.1);
11736            if logits.len() != (end - pos) * lm.1 || vocab == 0 {
11737                return MetalBatchNllOutcome::Failed(
11738                    "ordinary Metal NLL head returned an invalid shape".to_string(),
11739                );
11740            }
11741            for row in 0..(end - pos) {
11742                let absolute = pos + row;
11743                if absolute < start || absolute + 1 >= ids.len() {
11744                    continue;
11745                }
11746                let lg = &mut logits[row * lm.1..row * lm.1 + vocab];
11747                if let Some(mu) = self.logit_multiplier {
11748                    for v in lg.iter_mut() {
11749                        *v *= mu;
11750                    }
11751                }
11752                if let Some(c) = self.final_softcap {
11753                    for v in lg.iter_mut() {
11754                        *v = c * (*v / c).tanh();
11755                    }
11756                }
11757                let target = ids[absolute + 1] as usize;
11758                if target >= vocab {
11759                    return MetalBatchNllOutcome::Failed(format!(
11760                        "target token {target} exceeds Metal head rows {vocab}"
11761                    ));
11762                }
11763                let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
11764                let lse: f64 = lg
11765                    .iter()
11766                    .map(|&v| ((v - max) as f64).exp())
11767                    .sum::<f64>()
11768                    .ln()
11769                    + max as f64;
11770                nll += lse - lg[target] as f64;
11771                count += 1;
11772            }
11773            pos = end;
11774        }
11775        MetalBatchNllOutcome::Completed(nll, count)
11776    }
11777
11778    /// Commit a Metal verify round: replay the GDN recurrences over the
11779    /// `a + 1` accepted positions into the CPU states, append the accepted
11780    /// K/V rows from the mirrors to the CPU caches, re-point the mirrors.
11781    #[cfg(target_os = "macos")]
11782    fn metal_verify_commit(&mut self, a: usize) -> bool {
11783        let Some(mut pending) = self.metal_verify.take() else {
11784            return false;
11785        };
11786        let n = a + 1;
11787        // encode order == ascending layer order (the plan walks 0..layers)
11788        let idxs = pending.gdn_layers.clone();
11789        let mut outs: Vec<&mut [f32]> = self
11790            .kv_cache
11791            .layers
11792            .iter_mut()
11793            .enumerate()
11794            .filter(|(i, _)| idxs.binary_search(i).is_ok())
11795            .map(|(_, l)| l.linear_state.as_mut_slice())
11796            .collect();
11797        if !pending.graph.commit(n, &mut outs) {
11798            return false;
11799        }
11800        spec_stamp("c.replay");
11801        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
11802        // Read every layer before mutating any CPU cache.  Missing rows are
11803        // terminal after the replay has executed; never append a partial KV
11804        // prefix and continue on a serial path.
11805        let mut rows = Vec::with_capacity(pending.attn_layers.len());
11806        for (li, cpu_stored) in &pending.attn_layers {
11807            let mut kbuf = vec![0f32; n * nkv * hd];
11808            let mut vbuf = vec![0f32; n * nkv * hd];
11809            if !crate::gpu_metal::kv_mirror_read_rows(
11810                self.graph_kv_id,
11811                *li,
11812                nkv,
11813                hd,
11814                *cpu_stored,
11815                n,
11816                &mut kbuf,
11817                &mut vbuf,
11818            ) {
11819                return false;
11820            }
11821            rows.push((*li, *cpu_stored, kbuf, vbuf));
11822        }
11823        for (li, cpu_stored, kbuf, vbuf) in rows {
11824            let cache = &mut self.kv_cache.layers[li];
11825            for r in 0..n {
11826                cache.append(
11827                    &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
11828                    &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
11829                    &[],
11830                );
11831            }
11832            crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + n);
11833        }
11834        spec_stamp("c.kv");
11835        true
11836    }
11837
11838    /// The round's warm-ups as ONE b-row graph run over the MTP block on
11839    /// Metal: `pairs` = (trunk hidden, next token) at consecutive positions
11840    /// from `first_pos`; the block's input projection is folded in. This
11841    /// half encodes and SUBMITS (no wait); `mtp_warm_batch_finish` waits
11842    /// and pulls the appended K/V rows into the CPU MTP cache. None = the
11843    /// graph declined (nothing submitted, nothing appended).
11844    #[cfg(target_os = "macos")]
11845    fn mtp_warm_batch_submit(
11846        &mut self,
11847        m: &mut MtpModule,
11848        pairs: &[(&[f32], u32)],
11849        first_pos: usize,
11850    ) -> Option<MetalWarmPending> {
11851        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
11852        let b = pairs.len();
11853        if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
11854            return None;
11855        }
11856        let AttnKind::Full {
11857            wq,
11858            wk,
11859            wv,
11860            wo,
11861            q_norm,
11862            k_norm,
11863            output_gate,
11864            softplus_gate: None,
11865            bias: None,
11866        } = &m.layer.attn
11867        else {
11868            return None;
11869        };
11870        let FfnKind::Dense(d) = &m.layer.ffn else {
11871            return None;
11872        };
11873        if !d.segs.is_empty() {
11874            return None;
11875        }
11876        let (Some(pq), Some(pk), Some(pv), Some(po)) =
11877            (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
11878        else {
11879            return None;
11880        };
11881        let (Some(g), Some(u), Some(dn)) = (
11882            d.gate_proj.q1_parts(),
11883            d.up_proj.q1_parts(),
11884            d.down_proj.q1_parts(),
11885        ) else {
11886            return None;
11887        };
11888        let Some(eh) = m.eh_proj.q1_parts() else {
11889            return None;
11890        };
11891        let QTensor::Mapped { model, .. } = wq else {
11892            return None;
11893        };
11894        let model = model.clone();
11895        let hs = self.hidden_size;
11896        // [enorm(embed(tok)); hnorm(hidden)] rows
11897        let mut cat = vec![0f32; b * 2 * hs];
11898        for (j, (h, tok)) in pairs.iter().enumerate() {
11899            let e = self.embed_single(*tok);
11900            let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
11901            inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
11902            inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
11903        }
11904        let dims = GraphDims {
11905            hidden: hs,
11906            eps: self.rms_eps as f32,
11907            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11908        };
11909        spec_stamp("w.cat");
11910        let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
11911            return None;
11912        };
11913        spec_stamp("w.new");
11914        let l = AttnGpuLayer {
11915            attn_norm: &m.layer.input_norm,
11916            post_norm: &m.layer.post_norm,
11917            wq: pq,
11918            wk: pk,
11919            wv: pv,
11920            wo: po,
11921            ffn: MetalFfn::Dense {
11922                gate: g,
11923                up: u,
11924                down: dn,
11925            },
11926        };
11927        let (nh, nkv, hd, rd) = (
11928            self.num_heads,
11929            self.num_kv_heads,
11930            self.head_dim,
11931            self.rotary_dim,
11932        );
11933        let inv_freq = self.inv_freq.clone();
11934        let cpu_stored;
11935        {
11936            let cache = &m.kv;
11937            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11938            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11939            cpu_stored = cpu_k[0].len() / hd;
11940            // The cache may LAG the position (rows nobody warmed): the
11941            // pairs land at cpu_stored.. with their true RoPE positions
11942            // first_pos.., exactly what the one-by-one warm does. A cache
11943            // AHEAD of the position is a real inconsistency.
11944            if cpu_stored > first_pos {
11945                spec_stamp("w.decl");
11946                return None;
11947            }
11948            let p = AttnDeviceParams {
11949                kv_id: self.mtp_kv_id(),
11950                layer: Self::MTP_LAYER_BASE,
11951                nh,
11952                nkv,
11953                hd,
11954                rd,
11955                position: first_pos,
11956                scale: self.attn_scale,
11957                eps: self.rms_eps as f32,
11958                gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11959                late_qk_norm: self.qk_norm_after_rope,
11960                output_gate: *output_gate,
11961                q_norm: q_norm.as_deref(),
11962                k_norm: k_norm.as_deref(),
11963                inv_freq: &inv_freq,
11964                cpu_k,
11965                cpu_v,
11966                cpu_stored,
11967                o1: None,
11968            };
11969            if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
11970                return None;
11971            }
11972        }
11973        spec_stamp("w.enc");
11974        if !graph.submit() {
11975            return None;
11976        }
11977        spec_stamp("w.sub");
11978        Some(MetalWarmPending {
11979            graph,
11980            cpu_stored,
11981            b,
11982        })
11983    }
11984
11985    /// Submit and finish in one call (the prefill's MTP warm-up, where
11986    /// nothing runs in between).
11987    #[cfg(target_os = "macos")]
11988    fn mtp_warm_batch_metal(
11989        &mut self,
11990        m: &mut MtpModule,
11991        pairs: &[(&[f32], u32)],
11992        first_pos: usize,
11993    ) -> bool {
11994        match self.mtp_warm_batch_submit(m, pairs, first_pos) {
11995            Some(p) => self.mtp_warm_batch_finish(m, p),
11996            None => false,
11997        }
11998    }
11999
12000    /// Second half of the batched warm-up: wait for the submitted graph,
12001    /// pull its b appended K/V rows into the CPU MTP cache, re-point the
12002    /// mirror. False = the command buffer failed or the rows are missing
12003    /// (nothing appended; the caller falls back to the one-by-one warm).
12004    #[cfg(target_os = "macos")]
12005    fn mtp_warm_batch_finish(&mut self, m: &mut MtpModule, pending: MetalWarmPending) -> bool {
12006        let MetalWarmPending {
12007            mut graph,
12008            cpu_stored,
12009            b,
12010        } = pending;
12011        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12012        if !graph.sync() {
12013            return false;
12014        }
12015        spec_stamp("w.gpu");
12016        let mut kbuf = vec![0f32; b * nkv * hd];
12017        let mut vbuf = vec![0f32; b * nkv * hd];
12018        if !crate::gpu_metal::kv_mirror_read_rows(
12019            self.mtp_kv_id(),
12020            Self::MTP_LAYER_BASE,
12021            nkv,
12022            hd,
12023            cpu_stored,
12024            b,
12025            &mut kbuf,
12026            &mut vbuf,
12027        ) {
12028            return false;
12029        }
12030        for r in 0..b {
12031            m.kv.append(
12032                &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12033                &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12034                &[],
12035            );
12036        }
12037        crate::gpu_metal::kv_mirror_set_stored(
12038            self.mtp_kv_id(),
12039            Self::MTP_LAYER_BASE,
12040            cpu_stored + b,
12041        );
12042        spec_stamp("w.kv");
12043        true
12044    }
12045
12046    /// A committed token id from the high table (Cyrillic, CJK and the
12047    /// like sit above 131072 in Qwen's vocabulary; Latin subwords past
12048    /// the 65536 cut are rare enough to lose as rejected drafts) switches
12049    /// the draft to the full head for the next 16 tokens; other ids count
12050    /// down. On an M4 the full 660 MB head costs 5.5 ms a draft step
12051    /// against 1.4 for the shortlist, so the streak is kept short.
12052    pub(crate) fn note_draft_id(&mut self, id: u32) {
12053        let cut = Self::draft_vocab_rows(usize::MAX).max(131_072);
12054        if (id as usize) >= cut {
12055            self.draft_full_streak = 16;
12056        } else {
12057            self.draft_full_streak = self.draft_full_streak.saturating_sub(1);
12058        }
12059    }
12060
12061    /// The draft head's rows for the next step: the shortlist, or the full
12062    /// head while `draft_full_streak` runs.
12063    fn draft_head_rows(&self, head_rows: usize) -> usize {
12064        if self.draft_full_streak > 0 {
12065            head_rows
12066        } else {
12067            Self::draft_vocab_rows(head_rows)
12068        }
12069    }
12070
12071    /// Draft-head shortlist size: `CMF_DRAFT_VOCAB` rows (default 65536,
12072    /// capped at the head; 0 = full head).
12073    fn draft_vocab_rows(head_rows: usize) -> usize {
12074        static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
12075        let n = *N.get_or_init(|| {
12076            std::env::var("CMF_DRAFT_VOCAB")
12077                .ok()
12078                .and_then(|v| v.parse().ok())
12079                .unwrap_or(65536)
12080        });
12081        if n == 0 { head_rows } else { n.min(head_rows) }
12082    }
12083
12084    /// One MTP block step on the native Metal token graph: block input on
12085    /// the host, the attention layer + FFN device-resident over the MTP
12086    /// mirror, the head folded in when `want_logits`. The appended K/V row
12087    /// is pulled into the CPU MTP cache (owner of record) after the sync.
12088    #[cfg(target_os = "macos")]
12089    fn mtp_step_metal(
12090        &mut self,
12091        m: &mut MtpModule,
12092        hidden: &[f32],
12093        next_token: u32,
12094        position: usize,
12095        want_logits: bool,
12096    ) -> Option<(Vec<f32>, Vec<f32>)> {
12097        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12098        if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12099            || !crate::gpu::q1_force()
12100            || !crate::gpu::enabled_here()
12101            || self.attn_softcap > 0.0
12102            || self.attention_heads_per_layer.is_some()
12103            || m.kv.mode != crate::kv_cache::KvMode::F32
12104            || m.kv.o1.is_some()
12105        {
12106            return None;
12107        }
12108        let AttnKind::Full {
12109            wq,
12110            wk,
12111            wv,
12112            wo,
12113            q_norm,
12114            k_norm,
12115            output_gate,
12116            softplus_gate: None,
12117            bias: None,
12118        } = &m.layer.attn
12119        else {
12120            return None;
12121        };
12122        let FfnKind::Dense(d) = &m.layer.ffn else {
12123            return None;
12124        };
12125        if d.act != Act::Silu || !d.segs.is_empty() {
12126            return None;
12127        }
12128        let (pq, pk, pv, po) = (
12129            wq.q1_parts()?,
12130            wk.q1_parts()?,
12131            wv.q1_parts()?,
12132            wo.q1_parts()?,
12133        );
12134        let (g, u, dn) = (
12135            d.gate_proj.q1_parts()?,
12136            d.up_proj.q1_parts()?,
12137            d.down_proj.q1_parts()?,
12138        );
12139        let QTensor::Mapped { model, .. } = wq else {
12140            return None;
12141        };
12142        let model = model.clone();
12143        let lm = if want_logits {
12144            Some(self.weights.lm_head.q1_parts()?)
12145        } else {
12146            None
12147        };
12148        let dims = GraphDims {
12149            hidden: self.hidden_size,
12150            eps: self.rms_eps as f32,
12151            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12152        };
12153        // The block input `eh_proj · [enorm(e); hnorm(h)]` rides in the
12154        // graph (one submit a step); the host per-op matvec if it cannot.
12155        let hs = self.hidden_size;
12156        let mut x = vec![0f32; hs];
12157        let mut graph = TokenGraph::new(&model, dims, &x)?;
12158        let mut folded = false;
12159        if let Some(eh) = m.eh_proj.q1_parts() {
12160            let e = self.embed_single(next_token);
12161            let mut cat = vec![0.0f32; 2 * hs];
12162            let (cat_e, cat_h) = cat.split_at_mut(hs);
12163            inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
12164            inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
12165            folded = graph.encode_input_proj(eh, &cat);
12166        }
12167        if !folded {
12168            x = self.mtp_block_input(m, hidden, next_token);
12169            graph = TokenGraph::new(&model, dims, &x)?;
12170        }
12171        spec_stamp("d.in");
12172        let l = AttnGpuLayer {
12173            attn_norm: &m.layer.input_norm,
12174            post_norm: &m.layer.post_norm,
12175            wq: pq,
12176            wk: pk,
12177            wv: pv,
12178            wo: po,
12179            ffn: MetalFfn::Dense {
12180                gate: g,
12181                up: u,
12182                down: dn,
12183            },
12184        };
12185        let (nh, nkv, hd, rd) = (
12186            self.num_heads,
12187            self.num_kv_heads,
12188            self.head_dim,
12189            self.rotary_dim,
12190        );
12191        let inv_freq = self.inv_freq.clone();
12192        {
12193            let cache = &m.kv;
12194            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12195            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12196            let cpu_stored = cpu_k[0].len() / hd;
12197            let p = AttnDeviceParams {
12198                kv_id: self.mtp_kv_id(),
12199                layer: Self::MTP_LAYER_BASE,
12200                nh,
12201                nkv,
12202                hd,
12203                rd,
12204                position,
12205                scale: self.attn_scale,
12206                eps: self.rms_eps as f32,
12207                gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12208                late_qk_norm: self.qk_norm_after_rope,
12209                output_gate: *output_gate,
12210                q_norm: q_norm.as_deref(),
12211                k_norm: k_norm.as_deref(),
12212                inv_freq: &inv_freq,
12213                cpu_k,
12214                cpu_v,
12215                cpu_stored,
12216                o1: None,
12217            };
12218            if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12219                return None;
12220            }
12221        }
12222        // The draft's head over a vocabulary SHORTLIST (the first
12223        // CMF_DRAFT_VOCAB rows — BPE ids run roughly by merge rank, so the
12224        // low ids carry the mass): the verify keeps the full head, so a true
12225        // token past the cut is only a rejected draft, never a wrong token.
12226        // 662 MB a step on Qwen3.8 becomes 170 MB at 65536.
12227        let draft_rows = if let Some(lm) = lm {
12228            self.draft_head_rows(lm.1)
12229        } else {
12230            0
12231        };
12232        if let Some(lm) = lm {
12233            if !graph.lm_head_ok(lm) {
12234                return None;
12235            }
12236            if draft_rows < lm.1 {
12237                if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12238                    return None;
12239                }
12240            } else {
12241                graph.encode_lm_head(&m.final_norm, lm);
12242            }
12243        }
12244        spec_stamp("d.enc");
12245        if graph.sync_checked().is_err() {
12246            return None;
12247        }
12248        spec_stamp("d.gpu");
12249        let mut logits = Vec::new();
12250        if let Some(lm) = lm {
12251            let n_read = draft_rows.min(lm.1).min(self.vocab_size);
12252            logits = attention::take_buf(n_read);
12253            graph.read_logits(&mut logits);
12254            // ids past the shortlist: never drafted (−∞ in every chain)
12255            logits.resize(self.vocab_size, f32::NEG_INFINITY);
12256        }
12257        graph.finish(&mut x);
12258        let mut krow = attention::take_buf(nkv * hd);
12259        let mut vrow = attention::take_buf(nkv * hd);
12260        if crate::gpu_metal::kv_mirror_read_last(
12261            self.mtp_kv_id(),
12262            Self::MTP_LAYER_BASE,
12263            nkv,
12264            hd,
12265            &mut krow,
12266            &mut vrow,
12267        ) {
12268            m.kv.append(&krow, &vrow, &[]);
12269        }
12270        attention::recycle_buf(&mut krow);
12271        attention::recycle_buf(&mut vrow);
12272        spec_stamp("d.rd");
12273        Some((logits, x))
12274    }
12275
12276    /// `CMF_MTP_CHAIN=0` keeps the per-step draft (one submit and one
12277    /// host round trip per MTP step); the default drafts the whole chain
12278    /// in one command buffer when the round is plain greedy.
12279    ///
12280    /// Measured on an M4 (24 GB), Qwen3.8-27B q4tp, P3 at 160 tokens,
12281    /// k=7, six runs per arm alternating inside one lock window — the
12282    /// round's draft phase (median over the 34 rounds of a run) is
12283    /// 34.5 ms per round old against 30.1 new, i.e. 4.93 → 4.31 ms per
12284    /// draft step. That is the whole prize: the 7 submits cost ~0.6 ms
12285    /// each in host and submit latency and nothing else changes —
12286    /// acceptance (3.41 of 7) and tokens per round (4.41) are identical,
12287    /// and the round is 289 → 285 ms, decode 13.8 → 14.0 tok/s.
12288    fn mtp_chain_on() -> bool {
12289        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12290        *ON.get_or_init(|| std::env::var("CMF_MTP_CHAIN").as_deref() != Ok("0"))
12291    }
12292
12293    /// The round's k greedy drafts as ONE command buffer on Metal: the MTP
12294    /// block k times back to back, each step's token embedding gathered
12295    /// on the device from the argmax the step before it wrote, the head
12296    /// over the round's shortlist (or the full head during a full-head
12297    /// streak — decided once, before the chain, exactly as the per-step
12298    /// path decides it per step, since `draft_full_streak` only moves on
12299    /// a commit). One wait, then the k ids and the k appended K/V rows
12300    /// come back; the CPU MTP cache ends where k `mtp_step_metal` calls
12301    /// would have left it. `Err(false)` = declined before anything was
12302    /// committed (the per-step path takes the round); `Err(true)` = the
12303    /// command buffer failed after commit.
12304    #[cfg(target_os = "macos")]
12305    fn mtp_draft_chain_metal(
12306        &mut self,
12307        m: &mut MtpModule,
12308        hidden: &[f32],
12309        t_next: u32,
12310        position: usize,
12311        k: usize,
12312    ) -> Result<Vec<u32>, bool> {
12313        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12314        if k == 0
12315            || k > 64
12316            || !Self::mtp_chain_on()
12317            || std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12318            || !crate::gpu::q1_force()
12319            || !crate::gpu::enabled_here()
12320            || self.attn_softcap > 0.0
12321            || self.attention_heads_per_layer.is_some()
12322            || m.kv.mode != crate::kv_cache::KvMode::F32
12323            || m.kv.o1.is_some()
12324            // the chain gathers embeddings itself: only the plain table
12325            || self.dsv4.is_some()
12326            || self.dsv41.is_some()
12327            || self.qwen4_exp.is_some()
12328            || self.g3n.is_some()
12329        {
12330            return Err(false);
12331        }
12332        let AttnKind::Full {
12333            wq,
12334            wk,
12335            wv,
12336            wo,
12337            q_norm,
12338            k_norm,
12339            output_gate,
12340            softplus_gate: None,
12341            bias: None,
12342        } = &m.layer.attn
12343        else {
12344            return Err(false);
12345        };
12346        let FfnKind::Dense(d) = &m.layer.ffn else {
12347            return Err(false);
12348        };
12349        if d.act != Act::Silu || !d.segs.is_empty() {
12350            return Err(false);
12351        }
12352        let (Some(pq), Some(pk), Some(pv), Some(po)) =
12353            (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12354        else {
12355            return Err(false);
12356        };
12357        let (Some(g), Some(u), Some(dn)) = (
12358            d.gate_proj.q1_parts(),
12359            d.up_proj.q1_parts(),
12360            d.down_proj.q1_parts(),
12361        ) else {
12362            return Err(false);
12363        };
12364        let (Some(eh), Some(lm)) = (m.eh_proj.q1_parts(), self.weights.lm_head.q1_parts()) else {
12365            return Err(false);
12366        };
12367        let QTensor::Mapped { model, .. } = wq else {
12368            return Err(false);
12369        };
12370        let model = model.clone();
12371        // the embedding table: a q4tp tensor of the SAME blob, no Prism
12372        // inverse-embedding post-pass
12373        let QTensor::Mapped {
12374            model: em,
12375            idx: eidx,
12376            dtype: cortiq_core::TensorDtype::Q4TiledP,
12377            ..
12378        } = &self.weights.embed_tokens
12379        else {
12380            return Err(false);
12381        };
12382        if !std::sync::Arc::ptr_eq(em, &model)
12383            || crate::prism::is_inverse_embedding(&model, &model.tensors[*eidx].name)
12384        {
12385            return Err(false);
12386        }
12387        let embed = (
12388            *eidx,
12389            self.weights.embed_tokens.rows(),
12390            self.weights.embed_tokens.cols(),
12391        );
12392        if embed.2 != self.hidden_size || hidden.len() != self.hidden_size {
12393            return Err(false);
12394        }
12395        let dims = GraphDims {
12396            hidden: self.hidden_size,
12397            eps: self.rms_eps as f32,
12398            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12399        };
12400        let Some(mut graph) = TokenGraph::new(&model, dims, hidden) else {
12401            return Err(false);
12402        };
12403        if !graph.chain_embed_ok(embed) || !graph.lm_head_ok(lm) {
12404            return Err(false);
12405        }
12406        let l = AttnGpuLayer {
12407            attn_norm: &m.layer.input_norm,
12408            post_norm: &m.layer.post_norm,
12409            wq: pq,
12410            wk: pk,
12411            wv: pv,
12412            wo: po,
12413            ffn: MetalFfn::Dense {
12414                gate: g,
12415                up: u,
12416                down: dn,
12417            },
12418        };
12419        let (nh, nkv, hd, rd) = (
12420            self.num_heads,
12421            self.num_kv_heads,
12422            self.head_dim,
12423            self.rotary_dim,
12424        );
12425        let inv_freq = self.inv_freq.clone();
12426        let draft_rows = self.draft_head_rows(lm.1);
12427        let n_arg = draft_rows.min(lm.1).min(self.vocab_size);
12428        if n_arg == 0 {
12429            return Err(false);
12430        }
12431        // `CMF_MTP_CHAIN_SPLIT=1` commits each step as it is encoded, so
12432        // the GPU starts on step 0 while the host is still encoding step
12433        // 1 — a probe for whether the host encode is on the critical
12434        // path. It is not: three runs each, draft 30.0 ms per round split
12435        // against 30.1 whole, and the whole chain's host encode measures
12436        // 0.3 ms against a 29.7 ms wait. Kept as a probe, off by default.
12437        let split = std::env::var("CMF_MTP_CHAIN_SPLIT").as_deref() == Ok("1");
12438        let t_chain = std::time::Instant::now();
12439        graph.chain_ids_init(t_next, k);
12440        let cpu_stored;
12441        {
12442            let cache = &m.kv;
12443            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12444            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12445            cpu_stored = cpu_k[0].len() / hd;
12446            for j in 0..k {
12447                if !graph.encode_chain_input(
12448                    embed,
12449                    j as u32,
12450                    &m.enorm,
12451                    &m.hnorm,
12452                    self.embed_multiplier,
12453                    eh,
12454                ) {
12455                    return Err(false);
12456                }
12457                // step j's mirror row: the mirror is re-pointed at the CPU
12458                // rows before step 0 and advances by one per step; its
12459                // resync (never taken past step 0) reads the CPU rows
12460                let p = AttnDeviceParams {
12461                    kv_id: self.mtp_kv_id(),
12462                    layer: Self::MTP_LAYER_BASE,
12463                    nh,
12464                    nkv,
12465                    hd,
12466                    rd,
12467                    position: position + j,
12468                    scale: self.attn_scale,
12469                    eps: self.rms_eps as f32,
12470                    gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12471                    late_qk_norm: self.qk_norm_after_rope,
12472                    output_gate: *output_gate,
12473                    q_norm: q_norm.as_deref(),
12474                    k_norm: k_norm.as_deref(),
12475                    inv_freq: &inv_freq,
12476                    cpu_k: cpu_k.clone(),
12477                    cpu_v: cpu_v.clone(),
12478                    cpu_stored: cpu_stored + j,
12479                    o1: None,
12480                };
12481                if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12482                    return Err(false);
12483                }
12484                if draft_rows < lm.1 {
12485                    if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12486                        return Err(false);
12487                    }
12488                } else {
12489                    graph.encode_lm_head(&m.final_norm, lm);
12490                }
12491                if !graph.encode_argmax(n_arg, j as u32 + 1) {
12492                    return Err(false);
12493                }
12494                if split {
12495                    // CMF_MTP_CHAIN_SPLIT=1: commit every step so the GPU
12496                    // starts on step 0 while the host encodes the rest
12497                    graph.commit();
12498                }
12499            }
12500        }
12501        let t_enc = t_chain.elapsed();
12502        if graph.sync_checked().is_err() {
12503            return Err(true);
12504        }
12505        if std::env::var_os("CMF_GRAPH_SPEC_TIME").is_some() {
12506            eprintln!(
12507                "mtp-chain: encode {:.1} ms | wait {:.1} ms (k={k}, head rows {draft_rows}{})",
12508                t_enc.as_secs_f64() * 1e3,
12509                (t_chain.elapsed() - t_enc).as_secs_f64() * 1e3,
12510                if split { ", split" } else { "" }
12511            );
12512        }
12513        let mut ids = vec![0u32; k];
12514        if !graph.chain_ids_read(&mut ids) {
12515            return Err(true);
12516        }
12517        let mut kbuf = vec![0f32; k * nkv * hd];
12518        let mut vbuf = vec![0f32; k * nkv * hd];
12519        if !crate::gpu_metal::kv_mirror_read_rows(
12520            self.mtp_kv_id(),
12521            Self::MTP_LAYER_BASE,
12522            nkv,
12523            hd,
12524            cpu_stored,
12525            k,
12526            &mut kbuf,
12527            &mut vbuf,
12528        ) {
12529            return Err(true);
12530        }
12531        for r in 0..k {
12532            m.kv.append(
12533                &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12534                &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12535                &[],
12536            );
12537        }
12538        Ok(ids)
12539    }
12540
12541    fn try_batch_graph_wgpu(
12542        &self,
12543        hiddens: &mut [f32],
12544        positions: &[usize],
12545        k: usize,
12546        spec: Option<crate::gpu::SpecTail<'_>>,
12547    ) -> crate::gpu::BatchGraphOutcome {
12548        self.try_batch_graph_wgpu_prefix(hiddens, positions, k, spec, None)
12549    }
12550
12551    /// `try_batch_graph_wgpu` with the device-prefix mode: `layers_run`
12552    /// Some lets a stack that does not fit run its leading layers (the
12553    /// token graph's prefix rule) and reports how many; `hiddens` then
12554    /// holds the boundary rows and the caller runs the rest on the host.
12555    fn try_batch_graph_wgpu_prefix(
12556        &self,
12557        hiddens: &mut [f32],
12558        positions: &[usize],
12559        k: usize,
12560        spec: Option<crate::gpu::SpecTail<'_>>,
12561        layers_run: Option<&mut usize>,
12562    ) -> crate::gpu::BatchGraphOutcome {
12563        let graph_end = match self.mimo_moe.graph_prefix_end() {
12564            Some(end) if end < self.num_layers => {
12565                if layers_run.is_none() || spec.is_some() || end == 0 {
12566                    return crate::gpu::BatchGraphOutcome::Declined;
12567                }
12568                end
12569            }
12570            _ => self.num_layers,
12571        };
12572        let _tb = std::time::Instant::now();
12573        let batch_debug = std::env::var_os("CMF_BATCH_DEBUG").is_some();
12574        if self.attn_softcap > 0.0 {
12575            return crate::gpu::BatchGraphOutcome::Declined; // capped scores: no graph kernel — CPU path
12576        }
12577        // Same attention contract as the token graph: per-layer geometry
12578        // rides `geom`, anything it cannot express declines by name.
12579        if let Some(reason) = self.wgpu_graph_attn_decline() {
12580            self.note_graph_decline("wgpu batch graph", reason);
12581            return crate::gpu::BatchGraphOutcome::Declined;
12582        }
12583        let nh = self.num_heads;
12584        let (nkv, hd, rd) = self.layer_geom(0);
12585        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
12586        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
12587            if let Some((m, i, kind, rs)) = t
12588                .graph_weight()
12589                .or_else(|| t.graph_weight_descriptor())
12590            {
12591                let name = &m.tensors[i].name;
12592                let prism = if crate::prism::is_inverse_embedding(m, name) {
12593                    crate::gpu::GraphPrismOp::InverseEmbedding
12594                } else if crate::prism::is_forward_weight(m, name) {
12595                    crate::gpu::GraphPrismOp::Forward
12596                } else {
12597                    crate::gpu::GraphPrismOp::None
12598                };
12599                return Some(crate::gpu::GraphW {
12600                    idx: i,
12601                    kind,
12602                    row_scale: rs,
12603                    data: &[],
12604                    prism,
12605                    affine: crate::prism::is_affine_target(m, name),
12606                });
12607            }
12608            if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
12609                eprintln!(
12610                    "batch graph: tensor has no graph descriptor/f32 fallback rows={} cols={}",
12611                    t.rows(),
12612                    t.cols()
12613                );
12614            }
12615            t.as_f32().map(|d| crate::gpu::GraphW {
12616                idx: 0,
12617                kind: 4,
12618                row_scale: &[],
12619                data: d,
12620                prism: crate::gpu::GraphPrismOp::None,
12621                affine: false,
12622            })
12623        }
12624        let built: Option<(
12625            Vec<crate::gpu::GraphLayer<'_>>,
12626            std::sync::Arc<cortiq_core::CmfModel>,
12627        )> = (|| {
12628            let mut layers = Vec::with_capacity(graph_end);
12629            let mut model = None;
12630            for li in 0..graph_end {
12631                let lw = &self.weights.layers[self.phys_layer(li)];
12632                // MoE routes per token, so its experts are encoded token by
12633                // token inside the batched submit while attention and the
12634                // projections stay GEMMs. Refusing MoE here is what left
12635                // prefill running one position at a time: 33 tok/s against
12636                // 54 on decode, i.e. reading the prompt was slower than
12637                // writing the answer.
12638                let gffn = match &lw.ffn {
12639                    FfnKind::Dense(d) if !d.segs.is_empty() => {
12640                        if batch_debug {
12641                            eprintln!("batch graph: dense segmented FFN at layer {li}");
12642                        }
12643                        return None;
12644                    }
12645                    FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
12646                        gate: gw(&d.gate_proj)?,
12647                        up: gw(&d.up_proj)?,
12648                        down: gw(&d.down_proj)?,
12649                    },
12650                    FfnKind::Moe(m) => {
12651                        // Adaptive τ and expert masks stay on the CPU path.
12652                        // Sigmoid scores, the selection bias, a routed scale
12653                        // ≠ 1 and an ungated shared expert (hy_v3) ride the
12654                        // same flags word as the token graph — before, this
12655                        // refusal sent every Hy-MT2-30B prompt to the chunked
12656                        // fallback (8 tok/s of ingest against 53 of decode).
12657                        if m.route_tau.is_some() || m.mask.is_some() {
12658                            return None;
12659                        }
12660                        // A shared expert rides as slot top_k (gated or
12661                        // not is a flag on the select kernel); without one
12662                        // (MiMo-V2, LFM2-MoE) the kernels run top_k slots.
12663                        let shared = m.shared.as_ref();
12664                        let has_shared = shared.is_some();
12665                        let shared_gated = matches!(shared, Some((_, Some(_))));
12666                        let sgate = match shared {
12667                            Some((_, Some(sg))) => gw(sg)?,
12668                            // Ungated or absent: the router plane stands in
12669                            // so the plumbing stays total; the kernel pins
12670                            // weight 1 or never reads it.
12671                            _ => gw(&m.router)?,
12672                        };
12673                        let router = gw(&m.router)?;
12674                        // The batch MoE kernels still consume raw per-token
12675                        // rows and do not carry the descriptor-aware Prism
12676                        // transform/affine bit for router or shared-gate
12677                        // planes.  Refuse rather than route an untransformed
12678                        // source activation.
12679                        if router.prism != crate::gpu::GraphPrismOp::None
12680                            || router.affine
12681                            || sgate.prism != crate::gpu::GraphPrismOp::None
12682                            || sgate.affine
12683                        {
12684                            return None;
12685                        }
12686                        let inter = m.experts.first()?.gate_proj.rows();
12687                        let mut experts = Vec::with_capacity(m.experts.len() + 1);
12688                        let mut q4tp: Option<bool> = None;
12689                        let mut gu_q2: Option<bool> = None;
12690                        for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
12691                            if !matches!(e.act, Act::Silu)
12692                                || e.gate_proj.rows() != inter
12693                                || e.up_proj.rows() != inter
12694                            {
12695                                return None;
12696                            }
12697                            // Same ladder as the token graph: q4t → q2tp
12698                            // (mixed profile: 2-bit gate/up over a q4tp
12699                            // down) → q4tp. Uniform across the layer.
12700                            let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
12701                                Some((mm, gi)) => (
12702                                    mm,
12703                                    gi,
12704                                    e.up_proj.mapped_q4t()?.1,
12705                                    e.down_proj.mapped_q4t()?.1,
12706                                    false,
12707                                    false,
12708                                ),
12709                                None => match e.gate_proj.mapped_q2tp() {
12710                                    Some((mm, gi)) => (
12711                                        mm,
12712                                        gi,
12713                                        e.up_proj.mapped_q2tp()?.1,
12714                                        e.down_proj.mapped_q4tp()?.1,
12715                                        true,
12716                                        true,
12717                                    ),
12718                                    None => {
12719                                        let (mm, gi) = e.gate_proj.mapped_q4tp()?;
12720                                        (
12721                                            mm,
12722                                            gi,
12723                                            e.up_proj.mapped_q4tp()?.1,
12724                                            e.down_proj.mapped_q4tp()?.1,
12725                                            true,
12726                                            false,
12727                                        )
12728                                    }
12729                                },
12730                            };
12731                            if *q4tp.get_or_insert(is_p) != is_p
12732                                || *gu_q2.get_or_insert(is_q2) != is_q2
12733                            {
12734                                return None;
12735                            }
12736                            if [gi, ui, di].into_iter().any(|idx| {
12737                                mm.tensors
12738                                    .get(idx)
12739                                    .is_some_and(|t| {
12740                                        crate::prism::is_forward_weight(mm, &t.name)
12741                                            || crate::prism::is_affine_target(mm, &t.name)
12742                                    })
12743                            }) {
12744                                return None;
12745                            }
12746                            model.get_or_insert_with(|| mm.clone());
12747                            experts.push((gi, ui, di));
12748                        }
12749                        crate::gpu::GraphFfn::Moe {
12750                            router,
12751                            shared_gate: sgate,
12752                            experts,
12753                            n_exp: m.experts.len(),
12754                            top_k: m.top_k,
12755                            inter,
12756                            norm_topk: m.norm_topk_prob,
12757                            q4tp: q4tp?,
12758                            gu_q2: gu_q2.unwrap_or(false),
12759                            sigmoid: m.router_sigmoid,
12760                            bias: m.expert_bias.as_deref(),
12761                            has_shared,
12762                            shared_gated,
12763                            route_scale: m.routed_scaling,
12764                        }
12765                    }
12766                    _ => return None,
12767                };
12768                let attn = match &lw.attn {
12769                    AttnKind::Full {
12770                        wq,
12771                        wk,
12772                        wv,
12773                        wo,
12774                        q_norm,
12775                        k_norm,
12776                        output_gate,
12777                        softplus_gate,
12778                        bias,
12779                    } => {
12780                        if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
12781                            if batch_debug {
12782                                eprintln!(
12783                                    "batch graph: unsupported Full attention gate at layer {li} softplus={} heads={}",
12784                                    softplus_gate.is_some(),
12785                                    self.attention_heads_per_layer.is_some()
12786                                );
12787                            }
12788                            return None;
12789                        }
12790                        let (m, _, _, _) = wq
12791                            .graph_weight()
12792                            .or_else(|| wq.graph_weight_descriptor())?;
12793                        model = Some(m.clone());
12794                        crate::gpu::GraphAttn::Full {
12795                            wq: gw(wq)?,
12796                            wk: gw(wk)?,
12797                            wv: gw(wv)?,
12798                            wo: gw(wo)?,
12799                            q_norm: q_norm.as_deref(),
12800                            k_norm: k_norm.as_deref(),
12801                            late_qk_norm: self.qk_norm_after_rope,
12802                            bias: bias
12803                                .as_ref()
12804                                .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
12805                            output_gate: *output_gate,
12806                            cpu_k: self.kv_cache.layers[li].k_heads(),
12807                            cpu_v: self.kv_cache.layers[li].v_heads(),
12808                            geom: self.graph_attn_geom(li),
12809                        }
12810                    }
12811                    AttnKind::LinearGdn(w) => {
12812                        let Some(cfg) = self.gdn_cfg else {
12813                            if batch_debug {
12814                                eprintln!("batch graph: no GDN config at layer {li}");
12815                            }
12816                            return None;
12817                        };
12818                        let (m, _, _, _) = w
12819                            .in_proj_qkv
12820                            .graph_weight()
12821                            .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
12822                        model = Some(m.clone());
12823                        crate::gpu::GraphAttn::Gdn {
12824                            qkv: gw(&w.in_proj_qkv)?,
12825                            z: gw(&w.in_proj_z)?,
12826                            a: gw(&w.in_proj_a)?,
12827                            b: gw(&w.in_proj_b)?,
12828                            out: gw(&w.out_proj)?,
12829                            conv1d: &w.conv1d,
12830                            a_log: &w.a_log,
12831                            dt_bias: &w.dt_bias,
12832                            norm: &w.norm,
12833                            nv: cfg.num_v_heads,
12834                            nk: cfg.num_k_heads,
12835                            dk: cfg.key_head_dim,
12836                            dv: cfg.value_head_dim,
12837                            kk: cfg.conv_kernel,
12838                            cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
12839                        }
12840                    }
12841                    _ => return None,
12842                };
12843                layers.push(crate::gpu::GraphLayer {
12844                    input_norm: &lw.input_norm,
12845                    attn,
12846                    post_norm: &lw.post_norm,
12847                    ffn: gffn,
12848                });
12849            }
12850            Some((layers, model?))
12851        })();
12852        let Some((layers, model)) = built else {
12853            {
12854                use std::sync::atomic::{AtomicBool, Ordering};
12855                static SAID: AtomicBool = AtomicBool::new(false);
12856                if !SAID.swap(true, Ordering::Relaxed) {
12857                    tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
12858                }
12859            }
12860            return crate::gpu::BatchGraphOutcome::Declined;
12861        };
12862        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
12863            eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
12864        }
12865        crate::gpu::forward_batch_graph(
12866            &model,
12867            self.graph_kv_id,
12868            &layers,
12869            &self.inv_freq,
12870            hiddens,
12871            nh,
12872            nkv,
12873            hd,
12874            rd,
12875            self.hidden_size,
12876            self.intermediate_size,
12877            positions,
12878            self.kv_cache.max_seq_len,
12879            gemma,
12880            self.rms_eps as f32,
12881            self.attn_scale,
12882            k,
12883            &(0..graph_end)
12884                .map(|li| self.kv_cache.layers[self.phys_layer(li)].o1_views())
12885                .collect::<Vec<_>>(),
12886            self.o1_epoch,
12887            spec,
12888            layers_run,
12889        )
12890    }
12891
12892    /// Same, stopping after layer `upto` inclusive (routing probe φ).
12893    /// `CMF_DSV4_DRAFT_PROBE=1` — grade the draft against what the trunk goes on
12894    /// to produce. Off by default; it runs a whole draft per decoded token.
12895    fn draft_probe() -> bool {
12896        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12897        *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
12898    }
12899
12900    /// `CMF_DSV4_DRAFT_PROBE=1`: measure how much of the draft the trunk
12901    /// would have agreed with, WITHOUT verifying or rolling anything back.
12902    ///
12903    /// The number this produces decides the whole speculation design — at
12904    /// acceptance a, a block of B positions yields 1 + a + a² + ... tokens
12905    /// per trunk pass — so it is worth measuring before any of the machinery
12906    /// that would exploit it exists. Each draft is parked with the position
12907    /// it was made at, and graded as the real tokens arrive.
12908    /// `CMF_DSV4_SPEC=1` — the DeepSeek-V4 speculative decode: draft five
12909    /// on the card, verify them in one batched trunk pass, commit the
12910    /// accepted prefix, roll the rest back.
12911    #[cfg(feature = "gpu")]
12912    fn dsv4_spec_on() -> bool {
12913        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12914        *ON.get_or_init(|| {
12915            // Test-only runtime gate: model loading still performs the same
12916            // reservation and trunk packing, which gives rollback parity a
12917            // topology-identical non-speculative control arm.
12918            if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
12919                return v != "0";
12920            }
12921            // An explicit value is a diagnostic force/escape hatch.  With no
12922            // knob, speculation is eligible only when model loading reserved
12923            // its bounded pack.  On small q4tp cards the geometric reserve
12924            // gate deliberately leaves this at zero: trying to build DSpark
12925            // after the exact trunk filled VRAM is both slower and a device
12926            // OOM (measured on A40).
12927            std::env::var("CMF_DSV4_SPEC")
12928                .map(|v| v != "0")
12929                .unwrap_or_else(|_| {
12930                    crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
12931                })
12932        })
12933    }
12934
12935    /// One speculative round at the decode tip. `t_next` is the token the
12936    /// sampler just committed for `next_pos`. Returns the EXTRA accepted
12937    /// tokens (possibly none) and the new position, with `graph_logits`
12938    /// left holding the last accepted position's logits — exactly what the
12939    /// loop top expects. `None` means "speculate not this round": nothing
12940    /// was committed, the caller forwards normally.
12941    #[cfg(feature = "gpu")]
12942    fn dsv4_spec_step(
12943        &mut self,
12944        tip_token: u32,
12945        t_next: u32,
12946        next_pos: usize,
12947        max_extra: usize,
12948        drafted: &mut usize,
12949        accepted_ctr: &mut usize,
12950    ) -> Option<(Vec<u32>, usize)> {
12951        let t_all = std::time::Instant::now();
12952        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
12953            thread_local! {
12954                static LAST: std::cell::Cell<Option<std::time::Instant>> =
12955                    const { std::cell::Cell::new(None) };
12956            }
12957            LAST.with(|l| {
12958                if let Some(prev) = l.get() {
12959                    eprintln!(
12960                        "между раундами {:.1} мс",
12961                        prev.elapsed().as_secs_f64() * 1e3
12962                    );
12963                }
12964                l.set(Some(std::time::Instant::now()));
12965            });
12966        }
12967        if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
12968            eprintln!("spec_step: вход pos={next_pos}");
12969        }
12970        let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
12971        let cfg = self.dsv4.as_ref().map(|b| b.2)?;
12972        // The draft state and its capture, armed exactly as the probe does.
12973        if self.dspark.is_none() {
12974            let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
12975            if t.is_empty() {
12976                return None;
12977            }
12978            crate::dsv4::dspark_arm(&t, cfg.dim);
12979            self.dspark = Some(crate::dsv4::DsparkState::new(
12980                self.dsv4_mtp.len(),
12981                &cfg,
12982                t.len(),
12983            ));
12984        }
12985        let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
12986        let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
12987        if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
12988            eprintln!("spec_step: пак не построился (targets {targets:?})");
12989        }
12990        let pack = pack?;
12991        let block = crate::dsv4::dspark_block();
12992        let b_box = self.dsv4.as_mut()?;
12993        let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
12994        let ds = self.dspark.as_mut()?;
12995        // The tip's captures: either this token ran on a normal path that
12996        // filled the thread-local, or the previous spec round left them.
12997        let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
12998        if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
12999            if dbg {
13000                eprintln!("spec_step: нет захвата");
13001            }
13002            return None;
13003        }
13004        ds.have_hidden = true;
13005        let tip_pos = next_pos.checked_sub(1)?;
13006        let draft_started = std::time::Instant::now();
13007        let mut conf = Vec::new();
13008        let props = crate::dsv4::dspark_draft_gpu(
13009            g,
13010            &self.dsv4_mtp,
13011            &cfg,
13012            ds,
13013            pack,
13014            st.kv_id,
13015            tip_token,
13016            tip_pos,
13017            self.pool.as_deref(),
13018            &mut conf,
13019        );
13020        self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13021        *drafted += block;
13022        if props.is_empty() || props[0] != t_next {
13023            if dbg {
13024                eprintln!(
13025                    "spec_step: черновик {} (props0={:?} t_next={t_next})",
13026                    if props.is_empty() {
13027                        "пуст"
13028                    } else {
13029                        "мимо"
13030                    },
13031                    props.first()
13032                );
13033            }
13034            return None;
13035        }
13036        // `fed[0]` is `t_next`, which the outer loop has already committed;
13037        // only `fed[1..]` become additional output tokens. Cap the verify
13038        // transaction itself to the caller's remaining output budget instead
13039        // of merely truncating the returned vector: otherwise the KV/state
13040        // would advance past `max_tokens` and a 64-token request could return
13041        // 66 tokens (and poison a reused session with two invisible steps).
13042        let mut k_verify = crate::dsv4::dspark_verify_k()
13043            .min(props.len())
13044            .min(max_extra.saturating_add(1));
13045        // Adaptive depth: positions the draft itself doubts are paid for on
13046        // every verify and delivered almost never (natural-text survival
13047        // [.67 .50 .29 .08 .04]). `CMF_DSPARK_CONF_MIN=p` trims the fed
13048        // prefix at the first proposal whose confidence drops below p; on
13049        // predictable text the confidences stay high and nothing changes.
13050        let conf_min = {
13051            static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
13052            *M.get_or_init(|| {
13053                std::env::var("CMF_DSPARK_CONF_MIN")
13054                    .ok()
13055                    .and_then(|v| v.parse().ok())
13056                    .unwrap_or(0.0)
13057            })
13058        };
13059        if conf_min > 0.0 && conf.len() >= props.len() {
13060            let mut keep = 1usize;
13061            while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
13062                keep += 1;
13063            }
13064            k_verify = k_verify.min(keep.max(2));
13065        }
13066        if k_verify < 2 {
13067            return None;
13068        }
13069        let mut fed = Vec::with_capacity(k_verify);
13070        fed.push(t_next);
13071        fed.extend_from_slice(&props[1..k_verify]);
13072        let mut argmax = Vec::new();
13073        let mut logits_all = Vec::new();
13074        let mut walked = Vec::new();
13075        let txn = crate::dsv4::dsv4_verify_chunk(
13076            g,
13077            layers,
13078            &cfg,
13079            st,
13080            &fed,
13081            next_pos,
13082            &self.inv_freq,
13083            self.pool.as_deref(),
13084            &targets,
13085            &mut argmax,
13086            &mut logits_all,
13087            &mut walked,
13088        );
13089        if txn.is_none() && dbg {
13090            eprintln!("spec_step: verify отказал");
13091        }
13092        let txn = txn?;
13093        let spec_gpu_end = txn.gpu_end;
13094        let b = fed.len();
13095        let mut accepted = 1usize;
13096        while accepted < b && fed[accepted] == argmax[accepted - 1] {
13097            accepted += 1;
13098        }
13099        // `CMF_DSV4_SPEC_FORCE_REJECT=1` — accept nothing beyond the known
13100        // token, every round: the pure rollback exerciser. The output must
13101        // stay byte-identical to the plain walk; anything else is a
13102        // transaction bug, isolated from the acceptance logic.
13103        if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
13104            accepted = 1;
13105        }
13106        if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
13107            eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
13108        }
13109        let t_fin = std::time::Instant::now();
13110        if !crate::dsv4::dsv4_spec_finish(
13111            g,
13112            layers,
13113            &cfg,
13114            st,
13115            txn,
13116            accepted,
13117            &fed,
13118            &self.inv_freq,
13119            self.pool.as_deref(),
13120        ) {
13121            tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
13122            return None;
13123        }
13124        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13125            eprintln!(
13126                "finish(k={accepted}): {:.1} мс",
13127                t_fin.elapsed().as_secs_f64() * 1e3
13128            );
13129        }
13130        *accepted_ctr += accepted - 1;
13131        // Captures per accepted token: device targets photographed by the
13132        // batch, host targets from the verify's own walk. The last one
13133        // becomes the new tip's draft input; every one owes the ring an
13134        // entry for its position.
13135        let (hc, dim) = (cfg.hc_mult, cfg.dim);
13136        // Complete-chain layers are photographed by the fused submission;
13137        // partial device layers overwrite that slot after exact host cold-
13138        // expert correction.  Thus every target in the contiguous device
13139        // prefix has a valid per-token capture.
13140        let dev_caps: Vec<usize> = targets
13141            .iter()
13142            .copied()
13143            .filter(|&t| t < spec_gpu_end)
13144            .collect();
13145        let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
13146        if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
13147            return None;
13148        }
13149        for t in 0..accepted {
13150            let tip = t + 1 == accepted;
13151            for (slot, &tl) in targets.iter().enumerate() {
13152                if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
13153                    let lo = (di * b + t) * hc * dim;
13154                    crate::dsv4::dspark_capture(
13155                        &caps_all[lo..lo + hc * dim],
13156                        &cfg,
13157                        slot,
13158                        &mut ds.main_hidden,
13159                    );
13160                } else if tip
13161                    && crate::dsv4::dspark_peek_slot(slot, dim, {
13162                        let lo = slot * dim;
13163                        &mut ds.main_hidden[lo..lo + dim]
13164                    })
13165                {
13166                    // The tip's host-layer captures are the walk's own
13167                    // per-layer notes — exact. (The walk that ran last ended
13168                    // on exactly this token, on both the accept-all and the
13169                    // rollback path.)
13170                } else {
13171                    // Intermediate tokens: the post-tail state stands in for
13172                    // the per-layer capture on host targets below the last
13173                    // layer. Ring-entry quality only; the tip is exact.
13174                    crate::dsv4::dspark_capture(
13175                        &walked[t * hc * dim..(t + 1) * hc * dim],
13176                        &cfg,
13177                        slot,
13178                        &mut ds.main_hidden,
13179                    );
13180                }
13181            }
13182            crate::dsv4::dspark_ring_append(
13183                g,
13184                &self.dsv4_mtp,
13185                &cfg,
13186                ds,
13187                next_pos + t,
13188                self.pool.as_deref(),
13189            );
13190        }
13191        let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
13192        self.graph_logits = Some(row);
13193        // The speculative loop never runs the probe, so the trunk tally has
13194        // no other place to cycle. Armed only when someone asked for the
13195        // dump; the host tail is the only tallying path here, which is
13196        // precisely the population a partial pack would serve.
13197        if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
13198            crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
13199            crate::dsv4::pick_tally_arm();
13200        }
13201        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13202            eprintln!(
13203                "spec_step total {:.1} мс (k={accepted})",
13204                t_all.elapsed().as_secs_f64() * 1e3
13205            );
13206        }
13207        Some((fed[1..accepted].to_vec(), next_pos + accepted))
13208    }
13209
13210    fn dspark_probe(&mut self, position: usize, token_id: u32) {
13211        if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
13212            return;
13213        }
13214        // What the trunk just routed to, for this token.
13215        let trunk_now = crate::dsv4::pick_tally_take();
13216        crate::dsv4::trunk_freq_note(&trunk_now);
13217        if !trunk_now.is_empty() {
13218            self.dspark_trunk_picks.push(trunk_now);
13219            let keep = crate::dsv4::dspark_block();
13220            if self.dspark_trunk_picks.len() > keep {
13221                self.dspark_trunk_picks.remove(0);
13222            }
13223        }
13224        // Grade whatever is waiting: the token just decoded sits at
13225        // `position`, so it answers the draft made at `position - 1 - i`.
13226        for p in std::mem::take(&mut self.dspark_pending) {
13227            let Some(i) = position.checked_sub(p.0 + 1) else {
13228                continue;
13229            };
13230            let mut p = p;
13231            if i < p.1.len() {
13232                if p.2 && p.1[i] == token_id {
13233                    p.3 = i + 1;
13234                } else {
13235                    p.2 = false;
13236                }
13237                if i + 1 < p.1.len() {
13238                    self.dspark_pending.push(p);
13239                    continue;
13240                }
13241            }
13242            self.dspark_hist.push(p.3);
13243            self.dspark_real.push(token_id);
13244        }
13245        let Some(b) = &mut self.dsv4 else { return };
13246        let (g, layers, cfg) = (&b.0, &b.1, b.2);
13247        let n_layers = layers.len();
13248        if self.dspark.is_none() {
13249            let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13250            if t.is_empty() {
13251                return;
13252            }
13253            eprintln!(
13254                "DSpark: захват со слоёв {t:?}, блок {}",
13255                crate::dsv4::dspark_block()
13256            );
13257            crate::dsv4::dspark_arm(&t, cfg.dim);
13258            self.dspark = Some(crate::dsv4::DsparkState::new(
13259                self.dsv4_mtp.len(),
13260                &cfg,
13261                t.len(),
13262            ));
13263        }
13264        let ds = self.dspark.as_mut().unwrap();
13265        if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
13266            return; // this token ran on a path that captures nothing
13267        }
13268        let mut conf = Vec::new();
13269        crate::dsv4::pick_tally_arm();
13270        // The trunk has already consumed the adaptive VRAM budget. Until the
13271        // draft owns an explicit bounded device pack, its tensors are an
13272        // out-of-core CPU/disk tier by contract: never let per-op probes try
13273        // to squeeze another multi-gigabyte MTP expert cache onto the card.
13274        let draft_started = std::time::Instant::now();
13275        #[cfg(feature = "gpu")]
13276        let gpu_draft = crate::dsv4::dspark_gpu_on();
13277        #[cfg(not(feature = "gpu"))]
13278        let gpu_draft = false;
13279        let props = if gpu_draft {
13280            #[cfg(feature = "gpu")]
13281            {
13282                let kv_id = b.3.kv_id;
13283                match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
13284                    Some(pk) => crate::dsv4::dspark_draft_gpu(
13285                        g,
13286                        &self.dsv4_mtp,
13287                        &cfg,
13288                        ds,
13289                        pk,
13290                        kv_id,
13291                        token_id,
13292                        position,
13293                        self.pool.as_deref(),
13294                        &mut conf,
13295                    ),
13296                    None => Vec::new(),
13297                }
13298            }
13299            #[cfg(not(feature = "gpu"))]
13300            Vec::new()
13301        } else {
13302            crate::gpu::cpu_scope(|| {
13303                crate::dsv4::dspark_draft(
13304                    g,
13305                    &self.dsv4_mtp,
13306                    &cfg,
13307                    ds,
13308                    token_id,
13309                    position,
13310                    self.pool.as_deref(),
13311                    &mut conf,
13312                )
13313            })
13314        };
13315        self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13316        let draft_picks = crate::dsv4::pick_tally_take();
13317        crate::dsv4::dspark_freq_note(&draft_picks);
13318        // Re-arm for the NEXT trunk token; the probe runs after the forward,
13319        // so this is the only place that can.
13320        crate::dsv4::pick_tally_arm();
13321        if !props.is_empty() {
13322            // Two ratios, side by side: what a batched verify over the trunk
13323            // would read against what it asks for, and the same for the
13324            // draft's three stages. Near 1.0 means a batch amortises nothing.
13325            let (tu, tt) = {
13326                let flat: Vec<(usize, Vec<usize>)> = self
13327                    .dspark_trunk_picks
13328                    .iter()
13329                    .flat_map(|v| v.iter().cloned())
13330                    .collect();
13331                // Per layer, across the window of tokens.
13332                let mut per: std::collections::HashMap<usize, Vec<usize>> =
13333                    std::collections::HashMap::new();
13334                for (li, picks) in flat {
13335                    per.entry(li).or_default().extend(picks);
13336                }
13337                let n = per.len().max(1);
13338                let mut u = 0usize;
13339                let mut t = 0usize;
13340                for (_, v) in per {
13341                    t += v.len();
13342                    u += v.iter().collect::<std::collections::HashSet<_>>().len();
13343                }
13344                (u / n, t / n)
13345            };
13346            let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
13347            self.dspark_exp.push((tu, tt, du, dt));
13348            self.dspark_pending.push((position, props, true, 0));
13349        }
13350        if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
13351            let n = self.dspark_hist.len() as f32;
13352            let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
13353            let block = crate::dsv4::dspark_block();
13354            let mut at = vec![0usize; block + 1];
13355            for &k in &self.dspark_hist {
13356                at[k] += 1;
13357            }
13358            // Prefix survival: S_i = P(the first i positions all held).
13359            let mut surv = Vec::with_capacity(block);
13360            for i in 1..=block {
13361                let k = at[i..].iter().sum::<usize>() as f32 / n;
13362                surv.push(format!("{k:.2}"));
13363            }
13364            let distinct = self
13365                .dspark_real
13366                .iter()
13367                .collect::<std::collections::HashSet<_>>()
13368                .len();
13369            let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
13370                (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
13371            });
13372            let m = self.dspark_exp.len().max(1);
13373            eprintln!(
13374                "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
13375                 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
13376                self.dspark_hist.len(),
13377                mean + 1.0,
13378                surv.join(" ")
13379            );
13380            eprintln!(
13381                "DSpark: разных токенов {distinct} из {} (вырожденность), \
13382                 эксперты ствол {}/{} на слой за {block} токенов, \
13383                 черновик {}/{} за блок, draft {:.2} мс/блок",
13384                self.dspark_real.len(),
13385                tu / m,
13386                tt / m,
13387                du / m,
13388                dt / m,
13389                self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
13390            );
13391        }
13392    }
13393
13394    fn forward_layers_upto(
13395        &mut self,
13396        hidden: &[f32],
13397        position: usize,
13398        task_mask: Option<&TaskMask>,
13399        upto: Option<usize>,
13400    ) -> Vec<f32> {
13401        // In-process multi-GPU: each segment runs pinned to its card,
13402        // and the only thing crossing the boundary is one hidden vector
13403        // that never leaves this address space. Same layer split the
13404        // network mode does, minus the second process, the socket, the
13405        // serialization and the dir_hash handshake.
13406        if let Some(plan) = self.gpu_plan.clone() {
13407            if upto.is_none() && plan.len() > 1 {
13408                let mut h = hidden.to_vec();
13409                for &(dev, from, upto_incl) in plan.iter() {
13410                    h = crate::gpu::with_device(dev, || {
13411                        self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
13412                    });
13413                }
13414                return h;
13415            }
13416        }
13417        self.forward_layers_span(hidden, position, task_mask, 0, upto)
13418    }
13419
13420    /// Split this pipeline's layer stack across local GPUs: segment i
13421    /// runs on `devices[i]`. Contiguous and even by layer count — the
13422    /// VRAM-weighted planner is the next step, and an uneven card pair
13423    /// is why it will be needed. `None` clears the plan.
13424    pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
13425        self.set_gpu_plan_at(devices, None)
13426    }
13427
13428    /// The same, with an explicit first boundary (`--peer-split`): card
13429    /// 0 takes layers `[0..at)`, the rest split what remains. Uneven
13430    /// cards, or an attention-heavy head, are why this knob exists.
13431    pub fn set_gpu_plan_at(
13432        &mut self,
13433        devices: Option<&[usize]>,
13434        at: Option<usize>,
13435    ) -> Result<(), String> {
13436        let Some(devs) = devices.filter(|d| d.len() > 1) else {
13437            self.gpu_plan = None;
13438            return Ok(());
13439        };
13440        self.split_supported()?;
13441        let n = self.num_layers;
13442        if devs.len() > n {
13443            return Err(format!("{} devices for {n} layers", devs.len()));
13444        }
13445        if let Some(k) = at {
13446            if k == 0 || k >= n {
13447                return Err(format!("split at {k}: the model has {n} layers"));
13448            }
13449            if devs.len() == 2 {
13450                self.gpu_plan = Some(std::sync::Arc::new(vec![
13451                    (devs[0], 0, k - 1),
13452                    (devs[1], k, n - 1),
13453                ]));
13454                return Ok(());
13455            }
13456            return Err(format!(
13457                "an explicit split point takes exactly 2 devices, got {}",
13458                devs.len()
13459            ));
13460        }
13461        let per = n.div_ceil(devs.len());
13462        let mut plan = Vec::with_capacity(devs.len());
13463        let mut from = 0usize;
13464        for &d in devs {
13465            if from >= n {
13466                break;
13467            }
13468            let upto = (from + per - 1).min(n - 1);
13469            plan.push((d, from, upto));
13470            from = upto + 1;
13471        }
13472        self.gpu_plan = Some(std::sync::Arc::new(plan));
13473        Ok(())
13474    }
13475
13476    /// The active in-process split, if any: (device, first layer, last).
13477    pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
13478        self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
13479    }
13480
13481    /// Layer span [from ..= upto] (upto None = last layer): the building
13482    /// block the network pipeline-split rides on. `from > 0` skips the
13483    /// arch escape hatches (the pub `forward_span` refuses those archs
13484    /// first) and the whole-token graph — the plain per-layer loop is
13485    /// the canonical executor for a partial stack.
13486    fn embryo_qtensor_f32(t: &QTensor) -> Vec<f32> {
13487        if let Some(x) = t.as_f32() {
13488            return x.to_vec();
13489        }
13490        let mut out = vec![0.0; t.rows() * t.cols()];
13491        for r in 0..t.rows() {
13492            t.row_f32(r, &mut out[r * t.cols()..(r + 1) * t.cols()]);
13493        }
13494        out
13495    }
13496
13497    fn embryo_resident_eligible(&self) -> bool {
13498        // One mixer family per file: vmf_phase (kind 0/1) or
13499        // gated_delta_net (kind 4); the anchors are full (2) or bounded (3).
13500        if (self.vmf_cfg.is_none() && self.gdn_cfg.is_none())
13501            || self.num_layers != self.physical_layers
13502            || self.loop_final_norm
13503            || self.weights.layers.len() != self.num_layers
13504            || self.head_clusters.is_none()
13505            || self.final_softcap.is_some()
13506            || self.logit_multiplier.is_some()
13507            || self.attn_softcap != 0.0
13508            || self.mtp.is_some()
13509            || self.g3n.is_some()
13510            || self.dsv4.is_some()
13511            || self.dsv41.is_some()
13512            || self.qwen4_exp.is_some()
13513            // Dynamic routing swaps FFN weights mid-sequence under the
13514            // packed graph; a blend has no single overlay. A STATIC skill
13515            // (`from_model_with_skill`) is fine: the pack reads the live
13516            // `weights.layers[*].ffn`, i.e. the skill's tensors, and a
13517            // later `set_active_skill` change drops the pack
13518            // (`invalidate_for_weight_change`).
13519            || self.dyn_router.is_some()
13520            || self.dyn_phi_layer.is_some()
13521            || self.dyn_blend_loaded
13522            || self.o1_cfg.is_some()
13523            || self.swa.is_some()
13524            || self.sliding_layers.is_some()
13525            || self.global_attn.is_some()
13526            || self.attention_heads_per_layer.is_some()
13527            || self.attn_v_norm
13528            || self
13529                .kv_cache
13530                .layers
13531                .iter()
13532                .any(|l| l.mode != crate::kv_cache::KvMode::F32)
13533            || self.rope_scale != 1.0
13534            || self.rope_scale_local != 1.0
13535            || self.attn_scale != 1.0 / (self.head_dim as f32).sqrt()
13536            || self.hidden_size == 0
13537            || self.hidden_size > 1024
13538            || self.intermediate_size > 1024
13539            || self.num_heads == 0
13540            || self.num_kv_heads == 0
13541            || self.num_heads % self.num_kv_heads != 0
13542            || self.num_heads.saturating_mul(self.head_dim) > 1024
13543            || self.num_kv_heads.saturating_mul(self.head_dim) > 1024
13544            || self.vocab_size == 0
13545            || self.kv_cache.max_seq_len == 0
13546            || self.rotary_dim == 0
13547            || self.rotary_dim > self.head_dim
13548            || self.rotary_dim % 2 != 0
13549            || self.inv_freq.len() < self.rotary_dim / 2
13550        {
13551            return false;
13552        }
13553        // The resident shader is deliberately an f32 profile.  Dequantizing
13554        // a Q4/Q8 tensor into the packed buffer would silently change the
13555        // operator relative to the CPU quantized path, so quantized CMFs
13556        // retain the exact ordinary executor instead of claiming parity.
13557        // Measured consequence (RTX PRO 4000, S4 bounded export requantized
13558        // with `cortiq requant --quant q4tp-quantize`): `eligible=false`,
13559        // the generic wgpu whole-token graph refuses too, and the per-op
13560        // path decodes at ~73 tok/s against ~200 tok/s on the CPU q4tp
13561        // path — a q4tp Embryo-O1 file is a CPU artifact today; the
13562        // resident graph serves the f32 export.
13563        if self.weights.lm_head.as_f32().is_none()
13564            || self.weights.embed_tokens.as_f32().is_none()
13565            || self.weights.lm_head.rows() < self.vocab_size
13566            || self.weights.lm_head.cols() != self.hidden_size
13567            || self.weights.embed_tokens.rows() < self.vocab_size
13568            || self.weights.embed_tokens.cols() != self.hidden_size
13569            || self.weights.final_norm.len() != self.hidden_size
13570        {
13571            return false;
13572        }
13573        if let Some(cfg) = self.vmf_cfg {
13574            if cfg.num_heads.saturating_mul(cfg.nphase) > 1024
13575                || cfg.num_heads.saturating_mul(cfg.nphase.saturating_add(1)) > 1024
13576                || cfg.num_heads.saturating_mul(cfg.value_head_dim) > 1024
13577                || cfg.state_len() == 0
13578            {
13579                return false;
13580            }
13581        }
13582        if let Some(g) = self.gdn_cfg {
13583            // The resident GDN kernels (gpu_wgpu.rs `embryo_core_gdn_*`):
13584            // fused projection ≤ 2048 rows, nv·dv ≤ 1024, dk ≤ 128 lanes,
13585            // dv ≤ 256 lanes in vec4 rows, SiLU output gate (the Embryo
13586            // export), same rms eps as the stack.
13587            if g.num_v_heads == 0
13588                || g.num_k_heads == 0
13589                || g.num_v_heads % g.num_k_heads != 0
13590                || g.key_head_dim == 0
13591                || g.key_head_dim > 128
13592                || g.value_head_dim == 0
13593                || g.value_head_dim > 256
13594                || g.value_head_dim % 4 != 0
13595                || g.conv_kernel == 0
13596                || g.num_v_heads > 512
13597                || g.num_v_heads.saturating_mul(g.value_head_dim) > 1024
13598                || g.conv_dim() > 2048
13599                || g.conv_dim() % 4 != 0
13600                || g.hidden_size != self.hidden_size
13601                || g.output_gate_sigmoid
13602                || g.rms_eps != self.rms_eps
13603                || g.state_len() == 0
13604            {
13605                return false;
13606            }
13607        }
13608        let mut full_seen = false;
13609        for lw in &self.weights.layers {
13610            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
13611                return false;
13612            }
13613            match &lw.attn {
13614                AttnKind::LinearGdn(w) => {
13615                    let Some(g) = self.gdn_cfg else {
13616                        return false;
13617                    };
13618                    let (nv, dv, kk) = (g.num_v_heads, g.value_head_dim, g.conv_kernel);
13619                    if w.in_proj_qkv.rows() != g.conv_dim()
13620                        || w.in_proj_qkv.cols() != self.hidden_size
13621                        || w.in_proj_qkv.as_f32().is_none()
13622                        || w.in_proj_z.rows() != nv * dv
13623                        || w.in_proj_z.cols() != self.hidden_size
13624                        || w.in_proj_z.as_f32().is_none()
13625                        || w.in_proj_a.rows() != nv
13626                        || w.in_proj_a.cols() != self.hidden_size
13627                        || w.in_proj_a.as_f32().is_none()
13628                        || w.in_proj_b.rows() != nv
13629                        || w.in_proj_b.cols() != self.hidden_size
13630                        || w.in_proj_b.as_f32().is_none()
13631                        || w.conv1d.len() != g.conv_dim() * kk
13632                        || w.a_log.len() != nv
13633                        || w.dt_bias.len() != nv
13634                        || w.norm.len() != dv
13635                        || w.out_proj.rows() != self.hidden_size
13636                        || w.out_proj.cols() != nv * dv
13637                        || w.out_proj.as_f32().is_none()
13638                    {
13639                        return false;
13640                    }
13641                }
13642                AttnKind::Linear(w) => {
13643                    let Some(cfg) = self.vmf_cfg else {
13644                        return false;
13645                    };
13646                    if w.thq.rows() != cfg.num_heads * cfg.nphase
13647                        || w.thq.cols() != self.hidden_size
13648                        || w.thq.as_f32().is_none()
13649                        || w.thk.rows() != cfg.num_heads * cfg.nphase
13650                        || w.thk.cols() != self.hidden_size
13651                        || w.thk.as_f32().is_none()
13652                        || w.v_proj.rows() != cfg.num_heads * cfg.value_head_dim
13653                        || w.v_proj.cols() != self.hidden_size
13654                        || w.v_proj.as_f32().is_none()
13655                        || w.out_proj.rows() != self.hidden_size
13656                        || w.out_proj.cols() != cfg.num_heads * cfg.value_head_dim
13657                        || w.out_proj.as_f32().is_none()
13658                        || w.decay.len() != cfg.num_heads * 2 * cfg.nphase
13659                    {
13660                        return false;
13661                    }
13662                    if let Some((kg, kb)) = &w.k_gate {
13663                        if kg.rows() != cfg.num_heads
13664                            || kg.cols() != self.hidden_size
13665                            || kg.as_f32().is_none()
13666                            || kb.len() != cfg.num_heads
13667                        {
13668                            return false;
13669                        }
13670                    }
13671                    if let Some(conv) = &w.conv {
13672                        if conv.len() % self.hidden_size != 0 || conv.len() / self.hidden_size < 2 {
13673                            return false;
13674                        }
13675                    }
13676                }
13677                AttnKind::Full {
13678                    wq,
13679                    wk,
13680                    wv,
13681                    wo,
13682                    q_norm,
13683                    k_norm,
13684                    output_gate,
13685                    softplus_gate,
13686                    bias,
13687                } => {
13688                    if full_seen
13689                        || q_norm.is_some()
13690                        || k_norm.is_some()
13691                        || *output_gate
13692                        || softplus_gate.is_some()
13693                        || bias.is_some()
13694                        || wq.as_f32().is_none()
13695                        || wk.as_f32().is_none()
13696                        || wv.as_f32().is_none()
13697                        || wo.as_f32().is_none()
13698                        || wq.rows() != self.num_heads * self.head_dim
13699                        || wk.rows() != self.num_kv_heads * self.head_dim
13700                        || wv.rows() != self.num_kv_heads * self.head_dim
13701                        || wq.cols() != self.hidden_size
13702                        || wk.cols() != self.hidden_size
13703                        || wv.cols() != self.hidden_size
13704                        || wo.rows() != self.hidden_size
13705                        || wo.cols() != self.num_heads * self.head_dim
13706                    {
13707                        return false;
13708                    }
13709                    full_seen = true;
13710                }
13711                AttnKind::Bounded(w) => {
13712                    // The resident bounded attend scores S + W lanes in one
13713                    // 256-lane chunk; the format caps S + W at 160.
13714                    let Some(ac) = self.anchor_core.as_ref() else {
13715                        return false;
13716                    };
13717                    if self.bounded_rope.is_none()
13718                        || w.window != ac.window
13719                        || w.sink != ac.sink
13720                        || w.window == 0
13721                        || w.window + w.sink > 256
13722                        || w.sink_k.len() != self.num_kv_heads * w.sink * self.head_dim
13723                        || w.sink_v.len() != self.num_kv_heads * w.sink * self.head_dim
13724                        || w.wq.as_f32().is_none()
13725                        || w.wk.as_f32().is_none()
13726                        || w.wv.as_f32().is_none()
13727                        || w.wo.as_f32().is_none()
13728                        || w.wq.rows() != self.num_heads * self.head_dim
13729                        || w.wk.rows() != self.num_kv_heads * self.head_dim
13730                        || w.wv.rows() != self.num_kv_heads * self.head_dim
13731                        || w.wq.cols() != self.hidden_size
13732                        || w.wk.cols() != self.hidden_size
13733                        || w.wv.cols() != self.hidden_size
13734                        || w.wo.rows() != self.hidden_size
13735                        || w.wo.cols() != self.num_heads * self.head_dim
13736                    {
13737                        return false;
13738                    }
13739                }
13740                _ => return false,
13741            }
13742            match &lw.ffn {
13743                FfnKind::Dense(d) => {
13744                    if d.act != Act::Silu
13745                        || !d.segs.is_empty()
13746                        || d.gate_proj.as_f32().is_none()
13747                        || d.up_proj.as_f32().is_none()
13748                        || d.down_proj.as_f32().is_none()
13749                        || d.gate_proj.rows() != self.intermediate_size
13750                        || d.gate_proj.cols() != self.hidden_size
13751                        || d.up_proj.rows() != self.intermediate_size
13752                        || d.up_proj.cols() != self.hidden_size
13753                        || d.down_proj.rows() != self.hidden_size
13754                        || d.down_proj.cols() != self.intermediate_size
13755                    {
13756                        return false;
13757                    }
13758                }
13759                FfnKind::Moe(m) => {
13760                    if m.resonance.is_none()
13761                        || m.top_k != 1
13762                        || m.router_sigmoid
13763                        || !m.norm_topk_prob
13764                        || m.expert_bias.is_some()
13765                        || m.routed_scaling != 1.0
13766                        || m.route_tau.is_some()
13767                        || m.shared.is_none()
13768                        || m.mask.is_some()
13769                        || m.per_expert_scale.is_some()
13770                        || m.router_input_norm
13771                        || m.experts.is_empty()
13772                        || m.experts.len() > 8
13773                    {
13774                        return false;
13775                    }
13776                    let r = m.resonance.as_ref().unwrap();
13777                    if r.mu.len() != m.experts.len() * self.hidden_size
13778                        || r.bias.len() != m.experts.len()
13779                        || r.u.len() != m.experts.len() * r.k * self.hidden_size
13780                        || r.k > 128
13781                    {
13782                        return false;
13783                    }
13784                    let Some((shared, gate)) = &m.shared else {
13785                        return false;
13786                    };
13787                    if gate.is_some() || shared.act != Act::Silu || !shared.segs.is_empty() {
13788                        return false;
13789                    }
13790                    if shared.gate_proj.as_f32().is_none()
13791                        || shared.up_proj.as_f32().is_none()
13792                        || shared.down_proj.as_f32().is_none()
13793                        || shared.gate_proj.rows() != self.intermediate_size
13794                        || shared.gate_proj.cols() != self.hidden_size
13795                        || shared.up_proj.rows() != self.intermediate_size
13796                        || shared.up_proj.cols() != self.hidden_size
13797                        || shared.down_proj.rows() != self.hidden_size
13798                        || shared.down_proj.cols() != self.intermediate_size
13799                    {
13800                        return false;
13801                    }
13802                    for e in &m.experts {
13803                        if e.act != Act::Silu
13804                            || !e.segs.is_empty()
13805                            || e.gate_proj.as_f32().is_none()
13806                            || e.up_proj.as_f32().is_none()
13807                            || e.down_proj.as_f32().is_none()
13808                            || e.gate_proj.rows() != self.intermediate_size
13809                            || e.gate_proj.cols() != self.hidden_size
13810                            || e.up_proj.rows() != self.intermediate_size
13811                            || e.up_proj.cols() != self.hidden_size
13812                            || e.down_proj.rows() != self.hidden_size
13813                            || e.down_proj.cols() != self.intermediate_size
13814                        {
13815                            return false;
13816                        }
13817                    }
13818                }
13819                FfnKind::DenseMoe(_) => return false,
13820            }
13821            if lw.input_norm.len() != self.hidden_size || lw.post_norm.len() != self.hidden_size {
13822                return false;
13823            }
13824        }
13825        if full_seen && self.anchor_core.is_some() {
13826            return false;
13827        }
13828        full_seen || self.num_layers > 0
13829    }
13830
13831    /// The resident Embryo graph is the owner of this pipeline's forward:
13832    /// the same gate `forward_layers_span` applies before handing a token
13833    /// to `forward_embryo_graph` (both graph phases on, the explicit
13834    /// opt-in, a wgpu device, no earlier refusal, an eligible stack).
13835    fn embryo_resident_wanted(&self) -> bool {
13836        crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
13837            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
13838            && matches!(
13839                std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
13840                Ok("1") | Ok("parallel")
13841            )
13842            && crate::gpu::enabled_here()
13843            && !self.graph_refused()
13844            && self.embryo_resident_eligible()
13845    }
13846
13847    /// Chunked prefill on the resident graph: `ids` from `start` in
13848    /// chunks of `EMBRYO_CHUNK_MAX`, one submit each, projections/FFN as
13849    /// chunk GEMMs and the recurrent layers walked in time on the device.
13850    /// Returns the last position's logits when the whole span ran there.
13851    /// `None` = refused before any device work (the per-position path
13852    /// takes the span).  `CMF_EMBRYO_CHUNK=0` keeps the per-position
13853    /// prefill (A/B and the parity reference).
13854    fn embryo_prefill_chunked(&mut self, ids: &[u32], start: usize) -> Option<Vec<f32>> {
13855        if ids.len() < 2
13856            || std::env::var("CMF_EMBRYO_CHUNK").as_deref() == Ok("0")
13857            || !self.embryo_resident_wanted()
13858        {
13859            return None;
13860        }
13861        let model = self.ensure_embryo_graph()?;
13862        let cmax = std::env::var("CMF_EMBRYO_CHUNK")
13863            .ok()
13864            .and_then(|v| v.parse::<usize>().ok())
13865            .filter(|&v| v >= 1)
13866            .unwrap_or(crate::gpu::EMBRYO_CHUNK_MAX)
13867            .min(crate::gpu::EMBRYO_CHUNK_MAX);
13868        let hs = self.hidden_size;
13869        let n = ids.len();
13870        let mut pos = start;
13871        let mut last = None;
13872        let mut rows = Vec::with_capacity(cmax * hs);
13873        while pos < n {
13874            let end = (pos + cmax).min(n);
13875            rows.clear();
13876            for &id in &ids[pos..end] {
13877                rows.extend_from_slice(&self.embed_single(id));
13878            }
13879            let mut lg = Vec::new();
13880            if !crate::gpu::forward_embryo_graph_chunk(
13881                &model,
13882                self.graph_kv_id,
13883                &rows,
13884                pos,
13885                end - pos,
13886                &mut lg,
13887            ) {
13888                if pos == start {
13889                    if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
13890                        eprintln!("embryo-dbg: chunk prefill refused at position {pos}");
13891                    }
13892                    return None;
13893                }
13894                // The device sequence advanced through the earlier chunks;
13895                // a host continuation would mix two owners of the state.
13896                // Fail this sequence alone, leaving no stale state or key.
13897                self.kv_cache.clear();
13898                self.clear_history();
13899                crate::gpu::graph_kv_reset(self.graph_kv_id);
13900                panic!(
13901                    "resident Embryo chunk prefill refused at position {pos}; refusing an in-flight host fallback"
13902                );
13903            }
13904            last = Some(lg);
13905            pos = end;
13906        }
13907        if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
13908            eprintln!(
13909                "embryo-dbg: chunk prefill {} ids from {start} in {} submits (chunk {cmax})",
13910                n - start,
13911                (n - start).div_ceil(cmax)
13912            );
13913        }
13914        last
13915    }
13916
13917    fn ensure_embryo_graph(&mut self) -> Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>> {
13918        if self.embryo_graph.is_none() && self.embryo_resident_eligible() {
13919            const UMAX: u32 = u32::MAX;
13920            const HEADER: usize = crate::gpu::EMBRYO_META_HEADER;
13921            const REC: usize = 64;
13922            struct Pack {
13923                data: Vec<f32>,
13924            }
13925            impl Pack {
13926                fn put(&mut self, x: &[f32]) -> u32 {
13927                    if x.is_empty() {
13928                        return u32::MAX;
13929                    }
13930                    let off = self.data.len();
13931                    self.data.extend_from_slice(x);
13932                    off as u32
13933                }
13934            }
13935            // The mixer family of the file: vmf_phase geometry fills the
13936            // phase header words, gated_delta_net fills words 24..29.  A
13937            // file has exactly one linear core, so at most one is live.
13938            let vmf = self.vmf_cfg;
13939            let gdn = self.gdn_cfg;
13940            let mut pack = Pack { data: Vec::new() };
13941            let mut meta = vec![0u32; HEADER];
13942            meta[0] = self.hidden_size as u32;
13943            meta[1] = self.intermediate_size as u32;
13944            meta[2] = self.vocab_size as u32;
13945            meta[3] = self.num_layers as u32;
13946            meta[4] = vmf.map(|c| c.num_heads).unwrap_or(0) as u32;
13947            meta[5] = vmf.map(|c| c.nphase).unwrap_or(0) as u32;
13948            meta[6] = vmf.map(|c| c.value_head_dim).unwrap_or(0) as u32;
13949            if let Some(g) = gdn {
13950                meta[24] = g.num_v_heads as u32;
13951                meta[25] = g.num_k_heads as u32;
13952                meta[26] = g.key_head_dim as u32;
13953                meta[27] = g.value_head_dim as u32;
13954                meta[28] = g.conv_kernel as u32;
13955                meta[29] = g.conv_dim() as u32;
13956            }
13957            meta[7] = self.num_heads as u32;
13958            meta[8] = self.num_kv_heads as u32;
13959            meta[9] = self.head_dim as u32;
13960            meta[10] = self.kv_cache.max_seq_len as u32;
13961            let clusters = self.head_clusters.as_ref().unwrap();
13962            let cluster_count = clusters.len() / self.hidden_size;
13963            if clusters.len() % self.hidden_size != 0
13964                || cluster_count == 0
13965                || cluster_count > 1024
13966                || self.vocab_size % cluster_count != 0
13967                || self.weights.lm_head.rows() < self.vocab_size
13968                || self.weights.final_norm.len() != self.hidden_size
13969            {
13970                return None;
13971            }
13972            meta[11] = cluster_count as u32;
13973            meta[12] = (self.vocab_size / cluster_count) as u32;
13974            meta[13] = vmf.map(|c| c.state_len()).unwrap_or(0) as u32;
13975            meta[16] = self.rotary_dim as u32;
13976            meta[17] = matches!(self.norm_style, NormStyle::Gemma) as u32;
13977            meta[19] = (self.rms_eps as f32).to_bits();
13978            let max_conv = self
13979                .weights
13980                .layers
13981                .iter()
13982                .filter_map(|lw| match &lw.attn {
13983                    AttnKind::Linear(w) => w.conv.as_ref().map(|v| v.len() / self.hidden_size),
13984                    _ => None,
13985                })
13986                .max()
13987                .unwrap_or(1);
13988            // One state slot per recurrent layer: the phase state plus its
13989            // hidden-wide conv ring, or the GDN record `[conv ring | S]`
13990            // (`GdnCfg::state_len`).  A file carries one mixer family, so
13991            // the stride is exactly that family's record and the device
13992            // state buffer equals the header's recurrent bytes.
13993            let phase_stride = vmf
13994                .map(|c| c.state_len() + max_conv.saturating_sub(1) * self.hidden_size)
13995                .unwrap_or(0);
13996            let gdn_stride = gdn.map(|g| g.state_len()).unwrap_or(0);
13997            let state_stride = phase_stride.max(gdn_stride);
13998            // Bounded genome: the KV plane of an anchor is its ring
13999            // `[kvh][W][hd]` K + V, and only anchors own one.  Legacy full
14000            // anchors keep the `max_seq` planes indexed by layer.
14001            let bounded = self.anchor_core.clone();
14002            let (anchor_window, anchor_sink) = bounded
14003                .as_ref()
14004                .map(|ac| (ac.window, ac.sink))
14005                .unwrap_or((0, 0));
14006            let kv_stride = if bounded.is_some() {
14007                2usize
14008                    .saturating_mul(self.num_kv_heads)
14009                    .saturating_mul(anchor_window)
14010                    .saturating_mul(self.head_dim)
14011            } else {
14012                2usize
14013                    .saturating_mul(self.num_kv_heads)
14014                    .saturating_mul(self.kv_cache.max_seq_len)
14015                    .saturating_mul(self.head_dim)
14016            };
14017            meta[14] = state_stride as u32;
14018            meta[15] = kv_stride as u32;
14019            meta[18] = anchor_window as u32;
14020            meta[20] = anchor_sink as u32;
14021            meta[21] = match &self.bounded_rope {
14022                Some(rope) => {
14023                    // [W][half] cos then [W][half] sin, one contiguous table.
14024                    let off = pack.put(&rope.cos);
14025                    let _ = pack.put(&rope.sin);
14026                    off
14027                }
14028                None => UMAX,
14029            };
14030            let mut full_seen = false;
14031            let mut bounded_seen = 0usize;
14032            // Recurrent state slots belong to mixer layers only (phase or
14033            // GDN): an anchor owns a ring, not a state stride, so the
14034            // device state buffer is exactly the header's recurrent bytes.
14035            let mut phase_seen = 0usize;
14036            let mut gdn_seen = 0usize;
14037            for (li, lw) in self.weights.layers.iter().enumerate() {
14038                let base = meta.len();
14039                meta.resize(base + REC, UMAX);
14040                meta[base] = match &lw.attn {
14041                    AttnKind::Linear(w) if w.phase_delta => 1,
14042                    AttnKind::Linear(_) => 0,
14043                    AttnKind::Full { .. } => 2,
14044                    AttnKind::Bounded(_) => 3,
14045                    AttnKind::LinearGdn(_) => 4,
14046                    _ => UMAX,
14047                };
14048                meta[base + 1] = pack.put(&lw.input_norm);
14049                meta[base + 2] = pack.put(&lw.post_norm);
14050                meta[base + 25] = match &lw.attn {
14051                    AttnKind::Linear(_) => {
14052                        let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14053                        phase_seen += 1;
14054                        off
14055                    }
14056                    AttnKind::LinearGdn(_) => {
14057                        let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14058                        gdn_seen += 1;
14059                        off
14060                    }
14061                    _ => UMAX,
14062                };
14063                match &lw.attn {
14064                    AttnKind::LinearGdn(w) => {
14065                        // Layer record words 56..63 + 29, as the resident
14066                        // kernels read them (gpu_wgpu.rs `embryo_core_gdn_*`).
14067                        meta[base + 56] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_qkv));
14068                        meta[base + 57] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_z));
14069                        meta[base + 58] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_a));
14070                        meta[base + 59] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_b));
14071                        meta[base + 60] = pack.put(&w.conv1d);
14072                        meta[base + 61] = pack.put(&w.a_log);
14073                        meta[base + 62] = pack.put(&w.dt_bias);
14074                        meta[base + 63] = pack.put(&w.norm);
14075                        meta[base + 29] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14076                        meta[base + 24] = 0;
14077                    }
14078                    AttnKind::Linear(w) => {
14079                        meta[base + 3] = pack.put(&Self::embryo_qtensor_f32(&w.thq));
14080                        meta[base + 4] = pack.put(&Self::embryo_qtensor_f32(&w.thk));
14081                        meta[base + 5] = pack.put(&Self::embryo_qtensor_f32(&w.v_proj));
14082                        meta[base + 6] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14083                        let decay: Vec<f32> = w.decay.iter().map(|&x| x as f32).collect();
14084                        meta[base + 7] = pack.put(&decay);
14085                        if let Some((kg, kb)) = &w.k_gate {
14086                            meta[base + 8] = pack.put(&Self::embryo_qtensor_f32(kg));
14087                            meta[base + 9] = pack.put(kb);
14088                        }
14089                        if let Some(conv) = &w.conv {
14090                            meta[base + 10] = pack.put(conv);
14091                            meta[base + 24] = (conv.len() / self.hidden_size) as u32;
14092                        } else {
14093                            meta[base + 24] = 0;
14094                        }
14095                    }
14096                    AttnKind::Full { wq, wk, wv, wo, .. } => {
14097                        full_seen = true;
14098                        meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(wq));
14099                        meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(wk));
14100                        meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(wv));
14101                        meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(wo));
14102                        meta[base + 26] = (li * kv_stride) as u32;
14103                    }
14104                    AttnKind::Bounded(w) => {
14105                        meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(&w.wq));
14106                        meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(&w.wk));
14107                        meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(&w.wv));
14108                        meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(&w.wo));
14109                        // Ring slot of this anchor (anchors only, packed).
14110                        meta[base + 26] = (bounded_seen * kv_stride) as u32;
14111                        meta[base + 27] = pack.put(&w.sink_k);
14112                        meta[base + 28] = pack.put(&w.sink_v);
14113                        bounded_seen += 1;
14114                    }
14115                    _ => return None,
14116                }
14117                match &lw.ffn {
14118                    FfnKind::Dense(d) => {
14119                        meta[base + 15] = 0;
14120                        meta[base + 16] = 0;
14121                        meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&d.gate_proj));
14122                        meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&d.up_proj));
14123                        meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&d.down_proj));
14124                    }
14125                    FfnKind::Moe(m) => {
14126                        let r = m.resonance.as_ref().unwrap();
14127                        let (shared, _) = m.shared.as_ref().unwrap();
14128                        meta[base + 15] = 1;
14129                        meta[base + 16] = m.experts.len() as u32;
14130                        meta[base + 17] = pack.put(&r.mu);
14131                        meta[base + 18] = pack.put(&r.u);
14132                        meta[base + 19] = pack.put(&r.bias);
14133                        meta[base + 20] = r.k as u32;
14134                        // Word 30: the growth shell as the runtime applies
14135                        // it now (`+inf` on trunk rows, the stored finite
14136                        // shell on grown rows, all `+inf` under
14137                        // `CMF_GROWTH_SHELL=off`), followed by one `−∞`
14138                        // sentinel at index E the kernel writes as the
14139                        // score of an expert outside its shell — WGSL has
14140                        // no infinity literal, so the value travels as
14141                        // data (`embryo_core_route_finalize`).
14142                        let mut shell = r.effective_shell(m.experts.len());
14143                        shell.push(f32::NEG_INFINITY);
14144                        meta[base + 30] = pack.put(&shell);
14145                        meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&shared.gate_proj));
14146                        meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&shared.up_proj));
14147                        meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&shared.down_proj));
14148                        for (e, ex) in m.experts.iter().enumerate() {
14149                            meta[base + 32 + e * 3] =
14150                                pack.put(&Self::embryo_qtensor_f32(&ex.gate_proj));
14151                            meta[base + 33 + e * 3] =
14152                                pack.put(&Self::embryo_qtensor_f32(&ex.up_proj));
14153                            meta[base + 34 + e * 3] =
14154                                pack.put(&Self::embryo_qtensor_f32(&ex.down_proj));
14155                        }
14156                    }
14157                    FfnKind::DenseMoe(_) => return None,
14158                }
14159            }
14160            if !full_seen && self.num_layers == 0 {
14161                return None;
14162            }
14163            let id = {
14164                static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
14165                NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
14166            };
14167            let model = crate::gpu::EmbryoGraphModel {
14168                id,
14169                hidden: self.hidden_size,
14170                intermediate: self.intermediate_size,
14171                vocab: self.vocab_size,
14172                layers: self.num_layers,
14173                phase_heads: vmf.map(|c| c.num_heads).unwrap_or(0),
14174                nphase: vmf.map(|c| c.nphase).unwrap_or(0),
14175                phase_dv: vmf.map(|c| c.value_head_dim).unwrap_or(0),
14176                anchor_q_heads: self.num_heads,
14177                anchor_kv_heads: self.num_kv_heads,
14178                anchor_head_dim: self.head_dim,
14179                rotary_dim: self.rotary_dim,
14180                max_seq: self.kv_cache.max_seq_len,
14181                cluster_count,
14182                cluster_size: self.vocab_size / cluster_count,
14183                phase_state_len: vmf.map(|c| c.state_len()).unwrap_or(0),
14184                state_stride,
14185                kv_stride,
14186                norm_gemma: matches!(self.norm_style, NormStyle::Gemma),
14187                phase_mass: vmf.map(|c| c.phase_mass).unwrap_or(0.0),
14188                weights: pack.data,
14189                meta,
14190                lm_head: Self::embryo_qtensor_f32(&self.weights.lm_head),
14191                clusters: clusters.as_ref().clone(),
14192                final_norm: self.weights.final_norm.clone(),
14193                inv_freq: self.inv_freq.as_ref().clone(),
14194                bounded: bounded.is_some(),
14195                kv_layers: if bounded.is_some() {
14196                    bounded_seen
14197                } else {
14198                    self.num_layers
14199                },
14200                state_layers: phase_seen + gdn_seen,
14201                anchor_window,
14202                anchor_sink,
14203                phase_layers: phase_seen,
14204                gdn_layers: gdn_seen,
14205                gdn_heads: gdn.map(|g| g.num_v_heads).unwrap_or(0),
14206                gdn_k_heads: gdn.map(|g| g.num_k_heads).unwrap_or(0),
14207                gdn_dk: gdn.map(|g| g.key_head_dim).unwrap_or(0),
14208                gdn_dv: gdn.map(|g| g.value_head_dim).unwrap_or(0),
14209                gdn_kk: gdn.map(|g| g.conv_kernel).unwrap_or(0),
14210            };
14211            self.embryo_graph = Some(std::sync::Arc::new(model));
14212        }
14213        self.embryo_graph.clone()
14214    }
14215
14216    fn forward_layers_span(
14217        &mut self,
14218        hidden: &[f32],
14219        position: usize,
14220        task_mask: Option<&TaskMask>,
14221        from: usize,
14222        upto: Option<usize>,
14223    ) -> Vec<f32> {
14224        debug_assert!(
14225            from == 0
14226                || (self.dsv4.is_none()
14227                    && self.dsv41.is_none()
14228                    && self.qwen4_exp.is_none()
14229                    && self.g3n.is_none())
14230        );
14231        // Every plain forward — the whole-token Metal graph (`q1_graph_gpu`
14232        // wraps the GDN owners zero-copy and reallocates them on a size
14233        // change) and the CPU layer loop (reads/swaps `linear_state`) —
14234        // must see the previous speculative commit's asynchronous replay
14235        // complete. One mutex probe when nothing is pending.
14236        #[cfg(target_os = "macos")]
14237        if !crate::gpu_metal::wait_replay() {
14238            self.fail_metal_graph("the pending async replay failed before a plain forward");
14239            return vec![0.0; self.hidden_size];
14240        }
14241        if let Some(b) = &mut self.qwen4_exp {
14242            let _ = (task_mask, upto);
14243            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14244            let mut logits = Vec::new();
14245            crate::qwen4_exp::forward_token(
14246                &b.0,
14247                &b.1,
14248                &b.2,
14249                &mut b.3,
14250                token_id,
14251                position,
14252                &self.inv_freq,
14253                self.pool.as_deref(),
14254                &mut logits,
14255                true,
14256            );
14257            self.graph_logits = Some(logits);
14258            return vec![0.0; self.hidden_size];
14259        }
14260        // DeepSeek-V4 runs its own stack: the state is hc_mult copies, and
14261        // the forward returns LOGITS, not a hidden — the head is inside it
14262        // (the final fold sits between the last layer and the norm). The
14263        // token id rides in `hidden[0]`, written by embed_single, because
14264        // the hash layers route by id rather than by content.
14265        if let Some(b) = &mut self.dsv4 {
14266            let _ = (task_mask, upto);
14267            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14268            let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
14269            st.pos = position;
14270            let mut logits = Vec::new();
14271            crate::dsv4::forward_token(
14272                g,
14273                layers,
14274                &cfg,
14275                st,
14276                token_id,
14277                &self.inv_freq,
14278                self.pool.as_deref(),
14279                &mut logits,
14280            );
14281            self.graph_logits = Some(logits);
14282            self.dspark_probe(position, token_id);
14283            // The caller expects a hidden; the logits went out of band, as
14284            // with the fused lm_head path.
14285            return vec![0.0; self.hidden_size];
14286        }
14287        // DeepSeek-V4.1 owns its complete stack and emits logits out of band.
14288        if let Some(b) = &mut self.dsv41 {
14289            let _ = (task_mask, upto);
14290            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14291            let mut logits = Vec::new();
14292            crate::dsv41::forward_token(
14293                &b.0,
14294                &b.1,
14295                &b.2,
14296                &mut b.3,
14297                token_id,
14298                position,
14299                self.pool.as_deref(),
14300                &mut logits,
14301            );
14302            self.graph_logits = Some(logits);
14303            return vec![0.0; self.hidden_size];
14304        }
14305        // Gemma-3n runs its own stack (4 AltUp replicas don't fit this
14306        // loop); `hidden` is the extended embedding from embed_single.
14307        if let Some(b) = &self.g3n {
14308            let _ = (task_mask, upto);
14309            return crate::g3n::g3n_forward(
14310                &b.0,
14311                &b.1,
14312                hidden,
14313                position,
14314                &mut self.kv_cache.layers,
14315                self.num_heads,
14316                self.num_kv_heads,
14317                self.head_dim,
14318                self.pool.as_deref(),
14319            );
14320        }
14321        // Cortiq Embryo owns a separate resident graph: phase recurrent
14322        // state, resonance routing, the GQA anchor KV and hierarchical head
14323        // all execute in one Vulkan submit. It is limited to a complete
14324        // unmasked stack; spans and task masks retain the exact host path.
14325        // `CMF_EMBRYO_DBG=1` names the gate that keeps a token off the
14326        // resident graph — every refusal below is otherwise silent.
14327        if from == 0
14328            && upto.is_none()
14329            && task_mask.is_none()
14330            && self.anchor_core.is_some()
14331            && std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1")
14332        {
14333            static ONCE: std::sync::Once = std::sync::Once::new();
14334            ONCE.call_once(|| {
14335                eprintln!(
14336                    "embryo-dbg: graph_on decode={} prefill={} resident_env={:?} enabled_here={} \
14337                     unsupported={} eligible={}",
14338                    crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode),
14339                    crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill),
14340                    std::env::var("CMF_EMBRYO_RESIDENT").ok(),
14341                    crate::gpu::enabled_here(),
14342                    self.graph_refused(),
14343                    self.embryo_resident_eligible(),
14344                );
14345            });
14346        }
14347        if from == 0
14348            && upto.is_none()
14349            && task_mask.is_none()
14350            // Embryo's recurrent/KV state has no host import path.  Do not
14351            // seed it for a prefill-only graph and then silently decode from
14352            // an empty CPU cache; both phases must select the resident owner.
14353            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14354            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14355            // This whole-token Embryo path remains explicitly opt-in.
14356            // `CMF_GPU_WGPU_GRAPH=1` still enables the mature generic graph,
14357            // but must not silently select this model-specific resident path.
14358            && matches!(
14359                std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14360                Ok("1") | Ok("parallel")
14361            )
14362            && crate::gpu::enabled_here()
14363            && !self.graph_refused()
14364            // Sequence owner: past position zero the device continues only
14365            // a sequence it holds. A host-owned sequence (the graph refused
14366            // at its start, or its prefix was prefilled on the host) keeps
14367            // the host path to its end — never a device attempt at p > 0
14368            // over an empty device image.
14369            && (position == 0 || self.device_sequence_position().is_some())
14370            && self.embryo_resident_eligible()
14371            && let Some(model) = self.ensure_embryo_graph()
14372        {
14373            let mut lg = Vec::new();
14374            if crate::gpu::forward_embryo_graph(&model, self.graph_kv_id, hidden, position, &mut lg)
14375            {
14376                self.graph_logits = Some(lg);
14377                return vec![0.0; self.hidden_size];
14378            }
14379            if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14380                eprintln!("embryo-dbg: forward_embryo_graph refused at position {position}");
14381            }
14382            // The refusal is THIS pipeline's: falling through once is safe
14383            // at position zero (the host owns the sequence from here),
14384            // while a refusal at a later position of a device-owned
14385            // sequence would mix a host KV/state path with a partial
14386            // device sequence.
14387            self.mark_graph_refused();
14388            if position != 0 {
14389                // Fail this sequence alone and leave nothing stale behind:
14390                // no reuse key, no host or device state for the next
14391                // request on this slot to "extend".
14392                self.kv_cache.clear();
14393                self.clear_history();
14394                crate::gpu::graph_kv_reset(self.graph_kv_id);
14395                panic!(
14396                    "resident Embryo graph refused at position {position}; refusing an in-flight host fallback"
14397                );
14398            }
14399        }
14400        let mut h = hidden.to_vec();
14401        // MiMo-V2 expert placement: decided before the graph or the per-op
14402        // arena can claim the budget the expert bank needs.
14403        self.mimo_moe_prepare();
14404        let _mimo_q8 = self.mimo_moe.is_on()
14405            .then(crate::qtensor::enter_full_gpu_q8_scope);
14406        // Split borrows: copy scalars / clone handles so the per-layer
14407        // cfg does not hold `&self` while the KV cache is `&mut`.
14408        let (nh, _nkv, _hd, hs, _rd, eps) = (
14409            self.num_heads,
14410            self.num_kv_heads,
14411            self.head_dim,
14412            self.hidden_size,
14413            self.rotary_dim,
14414            self.rms_eps,
14415        );
14416        let pool = self.pool.clone();
14417        // Opt-in wgpu token-graph attention (discrete Vulkan/DX12): the whole
14418        // attention sub-block runs resident in one submit. Off by default.
14419        // Whole-token wgpu graph: eligibility + arbitration.
14420        //  - explicit CMF_GPU_WGPU_GRAPH forces it on/off;
14421        //  - discrete adapters (4090: decode 76 -> 137 tok/s) and GDN
14422        //    hybrids (recurrent state device-resident, no CPU twin to
14423        //    race) TRUST it;
14424        //  - integrated/mobile adapters RACE it against the normal path
14425        //    at generation granularity (gpu::graph_race_*) — tiled
14426        //    mobile GPUs can turn the ~300-dispatch graph into seconds
14427        //    per token, while a fast phone GPU keeps its win.
14428        let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
14429        let graph_on = match graph_env.as_deref() {
14430            Some("0") => false,
14431            Some("prefill") => false, // decode keeps the per-op path
14432            Some(_) => true,
14433            // Unset: same discrete-only default as every other graph
14434            // site. "Is the GPU on" used to stand in here — which made
14435            // the 0.2 tok/s whole-token graph race-eligible on mobile
14436            // adapters and cost 12-14× on first tokens (cmfmobile
14437            // TUNING.md); integrated GPUs keep the per-op probe path.
14438            None => crate::gpu::wgpu_graph_default(),
14439        };
14440        let graph_trusted =
14441            graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
14442        let race_eligible = graph_on
14443            && upto.is_none()
14444            && task_mask.is_none()
14445            && from == 0
14446            && !self.graph_refused();
14447        let mut tail_start = 0usize;
14448        if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
14449            let t_graph = std::time::Instant::now();
14450            let mut lg = Vec::new();
14451            let mut gl = 0usize;
14452            let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
14453            let declined = built.is_none();
14454            let built = match built {
14455                Some(Ok(hh)) => Some(hh),
14456                Some(Err(())) => {
14457                    // O(1) state was admitted before the device failure; the
14458                    // CPU mirrors are stale by construction.  Clear the whole
14459                    // sequence and stop rather than walking that stale state.
14460                    self.clear_sequence_state();
14461                    self.graph_failed
14462                        .store(true, std::sync::atomic::Ordering::Relaxed);
14463                    self.cancel
14464                        .store(true, std::sync::atomic::Ordering::Relaxed);
14465                    tracing::error!("token graph failed after admission; sequence state cleared");
14466                    return vec![0.0; self.hidden_size];
14467                }
14468                None => None,
14469            };
14470            // Past the transient guards (o1 still collecting, a softcap)
14471            // a refusal is about the weights and will never change —
14472            // remember it instead of walking every layer again next
14473            // token.
14474            if declined && !self.o1_active() && self.attn_softcap == 0.0 {
14475                self.mark_graph_refused();
14476            }
14477            graph_note(built.is_some(), gl, self.num_layers);
14478            if let Some(hh) = built {
14479                let dur = t_graph.elapsed();
14480                if std::env::var("CMF_GRAPH_PROF").is_ok() {
14481                    eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
14482                }
14483                if gl > 0 && gl < self.num_layers {
14484                    // Device prefix: the graph ran layers 0..gl and handed
14485                    // back the boundary hidden — the loop below owns the
14486                    // tail. The prefix layers' KV/state advanced on the
14487                    // device; the tail's advances on the host below. One
14488                    // boundary crossing per token.
14489                    h = hh;
14490                    tail_start = gl;
14491                } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
14492                    if !graph_trusted {
14493                        crate::gpu::graph_race_record(true, dur);
14494                    }
14495                    if !lg.is_empty() {
14496                        // Graph produced logits (final-norm + lm_head folded in) —
14497                        // pad/cap to vocab and hand them to the sampler directly.
14498                        lg.resize(self.vocab_size, 0.0);
14499                        if let Some(c) = self.final_softcap {
14500                            for l in lg.iter_mut() {
14501                                *l = c * (*l / c).tanh();
14502                            }
14503                        }
14504                        self.graph_logits = Some(lg);
14505                    }
14506                    return hh;
14507                }
14508                // Hopeless first graph token: discard it and fall through
14509                // to the normal path. Safe exactly here — the prompt KV is
14510                // still CPU-owned (chunked prefill), so recomputing this
14511                // position is exact; the mirror's extra row is never read
14512                // (the race just settled on the normal path).
14513            }
14514        }
14515        // KIMI-LINEAR HAS NO SPLIT BUG. The 2.6× reported from the
14516        // model rotation (12.2 tok/s on one card against 4.6 on two)
14517        // was a single measurement of a model whose arm arbitration is
14518        // borderline, and it did not survive repetition. Three runs an
14519        // arm, same binary, back to back:
14520        //   probe on : 1 GPU 9.5 / 5.7 / 5.9   2 GPU 7.8 / 13.0 / 13.3
14521        //   pinned   : 1 GPU 5.6 / 5.3 / 5.2   2 GPU 3.5 / 4.2 / 3.4
14522        // With the arms pinned the split costs about 1.45×, which is
14523        // what a layer split costs. With the probe free, TWO CARDS RUN
14524        // FASTER — because for this model the CPU arm wins some op
14525        // classes and the probe finds that.
14526        //
14527        // Two things do stand, and both are measured. The token graph
14528        // builds NOTHING here (`covered 0 of 14 layers [0..14)`), so
14529        // every layer walks per-op on either arm — that is where the
14530        // headroom is, not in the split. And this model's benchmark is
14531        // unusable without `CMF_GPU_PROBE=0`: the arbitration alone
14532        // moves it by more than 2×.
14533        //
14534        // Span runs (network split): the graph covers exactly [from..=upto]
14535        // — one submit per SEGMENT per token. No race: its state is global
14536        // and calibrated on full stacks, so spans take the graph only where
14537        // it is trusted by default (discrete adapters / CMF_GPU_WGPU_GRAPH).
14538        let span = from > 0 || upto.is_some();
14539        if span && graph_on && task_mask.is_none() && graph_trusted {
14540            let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
14541            let mut lg = Vec::new();
14542            let mut gl = 0usize;
14543            let span_res =
14544                self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
14545            let span_res = match span_res {
14546                Some(Ok(hh)) => Some(hh),
14547                Some(Err(())) => {
14548                    self.clear_sequence_state();
14549                    self.graph_failed
14550                        .store(true, std::sync::atomic::Ordering::Relaxed);
14551                    self.cancel
14552                        .store(true, std::sync::atomic::Ordering::Relaxed);
14553                    tracing::error!(
14554                        "span token graph failed after admission; sequence state cleared"
14555                    );
14556                    return vec![0.0; self.hidden_size];
14557                }
14558                None => None,
14559            };
14560            graph_note(span_res.is_some(), gl, upto_excl - from);
14561            if std::env::var("CMF_GPU_DEBUG").is_ok() {
14562                // How much of the span the graph actually covered. A
14563                // prefix of nothing means every layer walks per-op and
14564                // the split's extra cost is elsewhere.
14565                static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
14566                if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
14567                    eprintln!(
14568                        "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
14569                        upto_excl - from,
14570                        span_res.is_some()
14571                    );
14572                }
14573            }
14574            if let Some(hh) = span_res {
14575                if gl == upto_excl - from {
14576                    if !lg.is_empty() {
14577                        lg.resize(self.vocab_size, 0.0);
14578                        if let Some(c) = self.final_softcap {
14579                            for l in lg.iter_mut() {
14580                                *l = c * (*l / c).tanh();
14581                            }
14582                        }
14583                        self.graph_logits = Some(lg);
14584                    }
14585                    crate::gpu::set_layer(-1);
14586                    return hh;
14587                }
14588                // Partial device prefix of the span: CPU owns the tail.
14589                h = hh;
14590                tail_start = from + gl;
14591            }
14592        }
14593        // Layers the host is about to run whose device mirror moved ahead
14594        // of the host cache (a device prefix that shrank since the prompt,
14595        // a batched-prefill prefix longer than this token's): bring their
14596        // rows over first. One comparison per layer when nothing lags.
14597        let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
14598
14599        // A partial graph is an explicit GPU-prefix / CPU-tail split. Keep
14600        // the tail PURE host-side: letting its QTensor hooks re-enter the
14601        // residency arena streams every omitted layer through Vulkan and the
14602        // driver's freed-allocation cache can grow to the full model size
14603        // (25.4 GiB observed with a 14 GiB budget on Granite 30B Q8_2F).
14604        // With a MiMo expert bank the tail is not a whole-layer host
14605        // stream: its experts run from the bank (never the arena) and its
14606        // projections stay per-op on the device, which the bank's placement
14607        // left room for.
14608        let host_tail = tail_start > from;
14609        let _host_tail = (host_tail && !self.mimo_moe.is_on()).then(crate::gpu::enter_cpu_scope);
14610        let automatic_gpu_prefix = self.automatic_gpu_prefix();
14611
14612        let _prof_layers = crate::cpuprof::time(crate::cpuprof::Slot::Layers);
14613        #[cfg(target_os = "macos")]
14614        let mut gpu_skip_until = 0usize;
14615        for li in tail_start.max(from)..self.num_layers {
14616            let _capacity_tail = automatic_gpu_prefix
14617                .filter(|&prefix| li >= prefix && !self.mimo_moe.is_dynamic(li, host_tail))
14618                .map(|_| crate::gpu::enter_cpu_scope());
14619            crate::gpu::set_layer(li as i64); // layer-split GPU/CPU (CMF_GPU_LAYERS)
14620            if let Some(u) = upto {
14621                if li > u {
14622                    break;
14623                }
14624            }
14625            if let Some(mask) = task_mask {
14626                if !mask.layer_alive(li) {
14627                    continue; // dead layer: residual pass-through
14628                }
14629            }
14630            // Whole-block q1 token graph: a run of consecutive q1
14631            // layers — GDN and full attention — executes with one sync
14632            // per CPU attend instead of per op (macOS/Metal).
14633            #[cfg(target_os = "macos")]
14634            {
14635                if li < gpu_skip_until {
14636                    continue;
14637                }
14638                if task_mask.is_none() {
14639                    let end = self.q1_graph_gpu(li, upto, position, &mut h);
14640                    if self
14641                        .graph_failed
14642                        .load(std::sync::atomic::Ordering::Relaxed)
14643                    {
14644                        // The graph may have mutated device state before a
14645                        // command-buffer error. Never continue with a CPU
14646                        // tail or read a stale host mirror after admission.
14647                        return vec![0.0; self.hidden_size];
14648                    }
14649                    if end > li {
14650                        gpu_skip_until = end;
14651                        // Looped Transformer: the graph stopped at a loop
14652                        // boundary — apply final norm before the next iteration.
14653                        if self.is_loop_end(end - 1) && end < self.num_layers {
14654                            h = inference::rms_norm(
14655                                &h,
14656                                &self.weights.final_norm,
14657                                self.rms_eps,
14658                                self.norm_style,
14659                            );
14660                        }
14661                        continue;
14662                    }
14663                }
14664            }
14665
14666            if task_mask.is_none() {
14667                match self.mimo_graph_layer_rows(li, &mut h, &[position]) {
14668                    crate::gpu::BatchGraphOutcome::Completed => continue,
14669                    crate::gpu::BatchGraphOutcome::Failed => return vec![0.0; self.hidden_size],
14670                    crate::gpu::BatchGraphOutcome::Declined => {},
14671                }
14672            }
14673            #[cfg(feature = "gpu")]
14674            self.pull_lagging_host_kv(li, li + 1, position);
14675            let lw = &self.weights.layers[self.phys_layer(li)];
14676            if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
14677                if tp.parse::<usize>().ok() == Some(position) {
14678                    let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
14679                    eprintln!(
14680                        "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
14681                        h[0], h[1]
14682                    );
14683                }
14684            }
14685            // Norm into the pipeline scratch — the returning rms_norm
14686            // allocated twice per layer per token (roadmap §3 P0).
14687            let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
14688            inference::rms_norm_into(
14689                &h,
14690                &lw.input_norm,
14691                self.rms_eps,
14692                self.norm_style,
14693                &mut self.ws.n1,
14694            );
14695            drop(prof);
14696
14697            let attn_out = match &lw.attn {
14698                AttnKind::Mla(w) => {
14699                    let inv_freq_l = self.layer_inv_freq(li);
14700                    let rs = self.layer_rope_scale(li);
14701                    let eps = self.rms_eps;
14702                    let pool = self.pool.clone();
14703                    mla_attention(
14704                        w,
14705                        &self.ws.n1,
14706                        &mut self.kv_cache.layers[li],
14707                        position,
14708                        &inv_freq_l,
14709                        rs,
14710                        eps,
14711                        pool.as_deref(),
14712                    )
14713                }
14714                AttnKind::Linear(w) => {
14715                    let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
14716                    vmf_phase_forward(
14717                        &self.ws.n1,
14718                        w,
14719                        &cfg,
14720                        &mut self.kv_cache.layers[li].linear_state,
14721                        self.pool.as_deref(),
14722                    )
14723                }
14724                AttnKind::Kda(w) => {
14725                    let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
14726                    crate::linear_core::kda_forward(
14727                        &self.ws.n1,
14728                        w,
14729                        &cfg,
14730                        &mut self.kv_cache.layers[li].linear_state,
14731                        self.pool.as_deref(),
14732                    )
14733                }
14734                AttnKind::LinearGdn(w) => {
14735                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
14736                    gdn_forward(
14737                        &self.ws.n1,
14738                        w,
14739                        &cfg,
14740                        &mut self.kv_cache.layers[li].linear_state,
14741                        self.pool.as_deref(),
14742                    )
14743                }
14744                AttnKind::ShortConv(w) => {
14745                    let cfg = self
14746                        .short_conv_cfg
14747                        .expect("short-conv layer without short_conv_cfg");
14748                    short_conv_forward(
14749                        &self.ws.n1,
14750                        w,
14751                        &cfg,
14752                        &mut self.kv_cache.layers[li].linear_state,
14753                        self.pool.as_deref(),
14754                    )
14755                }
14756                AttnKind::Bounded(w) => {
14757                    // Natively bounded anchor: insert into the ring, attend
14758                    // over sinks ∪ window. No position, nothing appended.
14759                    let rope = self
14760                        .bounded_rope
14761                        .clone()
14762                        .expect("bounded layer without an installed rotation table");
14763                    let cfg = crate::bounded::BoundedAttnCfg {
14764                        num_heads: self.num_heads,
14765                        num_kv_heads: self.num_kv_heads,
14766                        head_dim: self.head_dim,
14767                        hidden_size: hs,
14768                        scale: self.attn_scale,
14769                        rope: &rope,
14770                        pool: pool.as_deref(),
14771                    };
14772                    crate::bounded::bounded_attention(
14773                        &self.ws.n1,
14774                        w,
14775                        &mut self.kv_cache.layers[li],
14776                        &cfg,
14777                    )
14778                }
14779                AttnKind::Full {
14780                    wq,
14781                    wk,
14782                    wv,
14783                    wo,
14784                    q_norm,
14785                    k_norm,
14786                    output_gate,
14787                    softplus_gate,
14788                    bias,
14789                } if self.kv_cache.layers[li].o1_sealed() => {
14790                    // O(1) override: decode on the sealed Nyström state
14791                    // instead of the growing KV cache.
14792                    let inv_freq_l = self.layer_inv_freq(li);
14793                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14794                    let cfg = QwenAttnCfg {
14795                        num_heads: self.layer_num_heads(li),
14796                        num_kv_heads: nkv_l,
14797                        head_dim: hd_l,
14798                        hidden_size: hs,
14799                        position,
14800                        inv_freq: &inv_freq_l,
14801                        rotary_dim: rd_l,
14802                        scale: self.attn_scale,
14803                        softcap: self.attn_softcap,
14804                        window: None,
14805                        v_norm: self.attn_v_norm,
14806                        qk_norm_after_rope: self.qk_norm_after_rope,
14807                        q_norm: q_norm.as_deref(),
14808                        k_norm: k_norm.as_deref(),
14809                        output_gate: *output_gate,
14810                        softplus_gate: softplus_gate
14811                            .as_ref()
14812                            .map(|(gate, per_head)| (gate, *per_head)),
14813                        rope_scale: self.layer_rope_scale(li),
14814                        bias: bias
14815                            .as_ref()
14816                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
14817                        rms_eps: eps,
14818                        norm_style: self.norm_style,
14819                        pool: pool.as_deref(),
14820                        v_head_dim: self.layer_v_dim(li),
14821                    };
14822                    attention::qwen_attention_nystrom(
14823                        &self.ws.n1,
14824                        wq,
14825                        wk,
14826                        wv,
14827                        wo,
14828                        &mut self.kv_cache.layers[li],
14829                        &cfg,
14830                    )
14831                }
14832                AttnKind::Full {
14833                    wq,
14834                    wk,
14835                    wv,
14836                    wo,
14837                    q_norm,
14838                    k_norm,
14839                    output_gate,
14840                    softplus_gate,
14841                    bias,
14842                } => 'attn: {
14843                    // wgpu token-graph attention (opt-in): whole sub-block in
14844                    // one submit, device K/V mirror. q1 only, no gate/bias/mask.
14845                    // Its kernel has no window, sink or narrow-V slot and
14846                    // one mirror geometry: such models stay on the CPU attend.
14847                    let dropin_reason =
14848                        graph_on.then(|| self.graph_attn_decline_reason()).flatten();
14849                    if let Some(reason) = dropin_reason {
14850                        self.note_graph_decline("wgpu attn dropin", reason);
14851                    }
14852                    if graph_on
14853                        && dropin_reason.is_none()
14854                        && !*output_gate
14855                        && softplus_gate.is_none()
14856                        && self.attention_heads_per_layer.is_none()
14857                        && bias.is_none()
14858                        && task_mask.is_none()
14859                    {
14860                        let inv_freq_l = self.layer_inv_freq(li);
14861                        let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14862                        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
14863                        if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
14864                            wq.mapped_q1(),
14865                            wk.mapped_q1(),
14866                            wv.mapped_q1(),
14867                            wo.mapped_q1(),
14868                        ) {
14869                            let gm = gm.clone();
14870                            let mut out = vec![0f32; hs];
14871                            let cache = &self.kv_cache.layers[li];
14872                            if crate::gpu::attn_dropin(
14873                                &gm,
14874                                self.graph_kv_id,
14875                                li,
14876                                &self.ws.n1,
14877                                qi,
14878                                ki,
14879                                vi,
14880                                oi,
14881                                q_norm.as_deref(),
14882                                k_norm.as_deref(),
14883                                self.qk_norm_after_rope,
14884                                &inv_freq_l,
14885                                nh,
14886                                nkv_l,
14887                                hd_l,
14888                                rd_l,
14889                                hs,
14890                                position,
14891                                self.kv_cache.max_seq_len,
14892                                gemma,
14893                                eps as f32,
14894                                cache.k_heads(),
14895                                cache.v_heads(),
14896                                &mut out,
14897                            ) {
14898                                break 'attn out;
14899                            }
14900                        }
14901                    }
14902                    let masked = task_mask
14903                        .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
14904                        .unwrap_or(false);
14905                    let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
14906                    // The masked kernel knows one pipeline-wide geometry and
14907                    // RoPE table, no window and no sink.
14908                    let plain = self.layer_attn_plain(li);
14909                    match (masked, f32_view) {
14910                        // Historical masked path (f32 slices; the loader
14911                        // keeps masked models in f32).
14912                        (true, (Some(q), Some(k), Some(v), Some(o))) if plain => {
14913                            let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
14914                            attention::multi_head_attention(
14915                                &self.ws.n1,
14916                                q,
14917                                k,
14918                                v,
14919                                o,
14920                                &mut self.kv_cache.layers[li],
14921                                self.num_heads,
14922                                self.num_kv_heads,
14923                                self.head_dim,
14924                                self.hidden_size,
14925                                position,
14926                                &active_heads,
14927                                &self.inv_freq,
14928                            )
14929                        }
14930                        (masked, _) => {
14931                            if masked {
14932                                tracing::warn!(
14933                                    "layer {li}: head mask on quantized weights or on a \
14934                                     window/sink/per-layer-geometry layer not supported \
14935                                     yet — executing dense"
14936                                );
14937                            }
14938                            let inv_freq_l = self.layer_inv_freq(li);
14939                            let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14940                            let cfg = QwenAttnCfg {
14941                                num_heads: self.layer_num_heads(li),
14942                                num_kv_heads: nkv_l,
14943                                head_dim: hd_l,
14944                                hidden_size: hs,
14945                                position,
14946                                inv_freq: &inv_freq_l,
14947                                rotary_dim: rd_l,
14948                                scale: self.attn_scale,
14949                                softcap: self.attn_softcap,
14950                                window: self.layer_window(li),
14951                                v_norm: self.attn_v_norm,
14952                                qk_norm_after_rope: self.qk_norm_after_rope,
14953                                q_norm: q_norm.as_deref(),
14954                                k_norm: k_norm.as_deref(),
14955                                output_gate: *output_gate,
14956                                softplus_gate: softplus_gate
14957                                    .as_ref()
14958                                    .map(|(gate, per_head)| (gate, *per_head)),
14959                                rope_scale: self.layer_rope_scale(li),
14960                                bias: bias
14961                                    .as_ref()
14962                                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
14963                                rms_eps: eps,
14964                                norm_style: self.norm_style,
14965                                pool: pool.as_deref(),
14966                                v_head_dim: self.layer_v_dim(li),
14967                            };
14968                            attention::qwen_attention(
14969                                &self.ws.n1,
14970                                wq,
14971                                wk,
14972                                wv,
14973                                wo,
14974                                &mut self.kv_cache.layers[li],
14975                                &cfg,
14976                            )
14977                        }
14978                    }
14979                }
14980            };
14981            // Gemma sandwich norm: normalize the attention branch before
14982            // it joins the residual stream.
14983            let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
14984                Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
14985                None => attn_out,
14986            };
14987            let lw = &self.weights.layers[self.phys_layer(li)];
14988            let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
14989            inference::add_rmsnorm_fused_into(
14990                &mut h,
14991                &attn_out,
14992                &lw.post_norm,
14993                self.rms_eps,
14994                self.norm_style,
14995                &mut self.ws.p1,
14996            );
14997            drop(prof);
14998            let mut attn_out = attn_out;
14999            attention::recycle_buf(&mut attn_out);
15000            let post_normed = &self.ws.p1;
15001
15002            let ffn_masked = task_mask
15003                .map(|m| m.ffn_active_count(li) < self.intermediate_size)
15004                .unwrap_or(false);
15005            // One masked dense CONTRACT, dispatched by cost. The
15006            // activation-zeroing arm (the batched sweep's, validated
15007            // against the replica to 0.8%) computes the FULL fused FFN
15008            // and zeroes the dead — right whenever most neurons live.
15009            // The sparse arm reads ONLY active rows and down columns —
15010            // per-row dots are slower per element than the fused kernel,
15011            // so it pays only once the mask is deep enough. The 0.5
15012            // crossover is first-principles (fused kernels run ~2x the
15013            // per-row dot throughput); a shallow specialist (95% alive)
15014            // stays fused, a --target-sparsity bake flips arms on its
15015            // own weight.
15016            let ffn_out = match (ffn_masked, &lw.ffn) {
15017                // A defragged tube layer answers its own mask: the core
15018                // always runs, each tube runs when its bit is on, and
15019                // the tubes that are off are never read from the mmap.
15020                (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
15021                    let row = task_mask
15022                        .and_then(|tm| tm.ffn_masks.get(li))
15023                        .map(|v| v.as_slice());
15024                    tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
15025                }
15026                (true, FfnKind::Dense(d)) => {
15027                    let tm = task_mask.unwrap();
15028                    let alive = tm.ffn_active_count(li);
15029                    let deep = alive * 2 <= self.intermediate_size;
15030                    if deep && d.down_proj.sparse_col_ok() && !d.gate_proj.has_prism_contract() {
15031                        let active = tm.ffn_active_indices(li);
15032                        sparse_ffn_quant(
15033                            d,
15034                            post_normed,
15035                            &active,
15036                            self.hidden_size,
15037                            self.pool.as_deref(),
15038                        )
15039                    } else if deep
15040                        && let (Some(g), Some(u), Some(dn)) = (
15041                            d.gate_proj.as_f32(),
15042                            d.up_proj.as_f32(),
15043                            d.down_proj.as_f32(),
15044                        )
15045                    {
15046                        let active = tm.ffn_active_indices(li);
15047                        inference::sparse_ffn_forward(
15048                            post_normed,
15049                            g,
15050                            u,
15051                            dn,
15052                            self.hidden_size,
15053                            self.intermediate_size,
15054                            &active,
15055                            self.pool.as_deref(),
15056                        )
15057                    } else {
15058                        let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
15059                        dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
15060                    }
15061                }
15062                (true, FfnKind::Moe(m)) => {
15063                    // MoE is sparse by expert selection; a task mask
15064                    // narrows the ROUTABLE set via its expert fields
15065                    // (spec §5) when it carries them.
15066                    let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
15067                    ffn_forward(
15068                        &lw.ffn,
15069                        post_normed,
15070                        self.pool.as_deref(),
15071                        allowed.as_deref(),
15072                    )
15073                }
15074                (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
15075                    dm,
15076                    post_normed,
15077                    &h,
15078                    self.rms_eps,
15079                    self.norm_style,
15080                    self.pool.as_deref(),
15081                ),
15082                (false, _) => match &lw.ffn {
15083                    FfnKind::DenseMoe(dm) => dense_moe_ffn(
15084                        dm,
15085                        post_normed,
15086                        &h,
15087                        self.rms_eps,
15088                        self.norm_style,
15089                        self.pool.as_deref(),
15090                    ),
15091                    FfnKind::Moe(m)
15092                        if task_mask.is_none() && self.mimo_moe.is_dynamic(li, host_tail) =>
15093                    {
15094                        moe_ffn_banked(&mut self.mimo_moe, li, m, post_normed, self.pool.as_deref())
15095                    }
15096                    _ => {
15097                        let allowed = match (&lw.ffn, task_mask) {
15098                            (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
15099                            _ => None,
15100                        };
15101                        ffn_forward(
15102                            &lw.ffn,
15103                            post_normed,
15104                            self.pool.as_deref(),
15105                            allowed.as_deref(),
15106                        )
15107                    }
15108                },
15109            };
15110            let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
15111                Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
15112                None => ffn_out,
15113            };
15114            for (i, &f) in ffn_out.iter().enumerate() {
15115                h[i] += f;
15116            }
15117            let mut ffn_out = ffn_out;
15118            attention::recycle_buf(&mut ffn_out);
15119
15120            // Gemma-4: the layer output is scaled by a learned scalar.
15121            if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
15122                for v in h.iter_mut() {
15123                    *v *= sc;
15124                }
15125            }
15126            // CMF_LAYER_DUMP: this position's hidden after layer li.
15127            if self.layer_dump.is_some() {
15128                self.dump_layer_row(position, li, &h);
15129            }
15130
15131            // Looped Transformer: apply final norm at the end of each loop iteration.
15132            // Nanbeige 4.2: after layer 21 (virtual), apply norm before looping back to layer 0.
15133            if self.is_loop_end(li) && li + 1 < self.num_layers {
15134                h = inference::rms_norm(
15135                    &h,
15136                    &self.weights.final_norm,
15137                    self.rms_eps,
15138                    self.norm_style,
15139                );
15140            }
15141
15142            // Dynamic routing φ capture (on-policy): the
15143            // EMA of the post-residual hidden at the router's phi_layer,
15144            // updated as the context evolves during decode.
15145            if self.dyn_phi_layer == Some(li) {
15146                self.update_dyn_phi(&h);
15147            }
15148        }
15149        crate::gpu::set_layer(-1); // layers done — lm_head outside layer-split
15150        if let Some(t) = t_race_cpu {
15151            crate::gpu::graph_race_record(false, t.elapsed());
15152        }
15153
15154        h
15155    }
15156
15157    /// EMA of φ at the router layer (rolling, weight 0.2 = ~5-token
15158    /// horizon). First observation seeds it exactly.
15159    fn update_dyn_phi(&mut self, h: &[f32]) {
15160        const A: f32 = 0.2;
15161        if self.dyn_phi_ema.len() != h.len() {
15162            self.dyn_phi_ema = vec![0.0; h.len()];
15163            self.dyn_phi_seen = 0;
15164        }
15165        if self.dyn_phi_seen == 0 {
15166            self.dyn_phi_ema.copy_from_slice(h);
15167        } else {
15168            for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
15169                *e = (1.0 - A) * *e + A * v;
15170            }
15171        }
15172        self.dyn_phi_seen += 1;
15173    }
15174
15175    /// Current router φ (EMA at phi_layer); empty until first capture.
15176    pub fn dyn_phi(&self) -> &[f32] {
15177        &self.dyn_phi_ema
15178    }
15179
15180    /// Enable/disable φ capture at the router layer, reset the EMA.
15181    pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
15182        self.dyn_phi_layer = layer;
15183        self.dyn_phi_ema.clear();
15184        self.dyn_phi_seen = 0;
15185    }
15186
15187    /// Skills eligible for dynamic switching: (index, id, phi_layer).
15188    pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
15189        let Some(model) = &self.model else {
15190            return Vec::new();
15191        };
15192        model
15193            .header
15194            .skills
15195            .iter()
15196            .enumerate()
15197            .filter_map(|(i, sk)| {
15198                // A v2 record routes only through the request-level
15199                // backbone-gated decision (its status and gate are
15200                // checked there), never per token.
15201                if sk.is_v2() {
15202                    return None;
15203                }
15204                let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
15205                let sel = sk.selection.as_ref()?;
15206                (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
15207            })
15208            .collect()
15209    }
15210
15211    /// Index of the currently overlaid skill (None = backbone).
15212    pub fn active_skill(&self) -> Option<usize> {
15213        self.dyn_active
15214    }
15215
15216    /// Enable dynamic per-token skill routing: build the hysteresis
15217    /// router from the container's routable skills, start φ capture at
15218    /// their (shared) phi_layer. Returns the number of routable skills
15219    /// (0 = nothing to route; router stays off). Idempotent.
15220    pub fn enable_dynamic_routing(&mut self) -> usize {
15221        use crate::swarm::{DynRouter, RoutableSkill};
15222        let Some(model) = self.model.clone() else {
15223            return 0;
15224        };
15225        // Router policy v2 (spec §9.4) routes per REQUEST: the backbone is
15226        // the default and only the backbone-gated decision may pick a
15227        // skill. A per-token switch would bypass that gate (and change
15228        // the O(1) state mid-sequence), so the hysteresis router never
15229        // runs on such a file; the caller keeps the request-level
15230        // decision.
15231        if let Some(r) = &model.header.router {
15232            tracing::warn!(
15233                "dynamic routing disabled: this file declares router policy '{}' with \
15234                 granularity \"{}\" — the request-level decision applies instead",
15235                r.policy,
15236                r.granularity
15237            );
15238            return 0;
15239        }
15240        // Format-v2 skill records (bit SKILLS_V2) without a router policy:
15241        // their status/gate contract ("auto-routing requires active +
15242        // measured") lives in the backbone-gated decision only — the
15243        // hysteresis router would switch into a quarantined record
15244        // (fail-open). Refuse the whole file, not just its v2 records.
15245        if model.required_features & cortiq_core::format::features::SKILLS_V2 != 0
15246            || model.header.skills.iter().any(|s| s.is_v2())
15247        {
15248            tracing::warn!(
15249                "dynamic routing disabled: this file carries format-v2 skill records \
15250                 (SKILLS_V2) — they route per request through a router policy only"
15251            );
15252            return 0;
15253        }
15254        // A blend materialized f32 working tensors into the layers; there
15255        // is no single skill index to revert from → refuse (honest).
15256        if self.dyn_blend_loaded {
15257            tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
15258            return 0;
15259        }
15260        // A statically-overlaid skill that is NOT FFN-eligible can't be
15261        // cheaply reverted at generation start → refuse rather than
15262        // silently keep it overlaid.
15263        if let Some(a) = self.dyn_active {
15264            if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
15265                tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
15266                return 0;
15267            }
15268        }
15269        let hidden = self.hidden_size;
15270        let mut skills = Vec::new();
15271        for (idx, id, _phi) in self.dynamic_skills() {
15272            if let Some(sel) = model.header.skills[idx].selection.as_ref() {
15273                if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
15274                    skills.push(rs);
15275                }
15276            }
15277        }
15278        if skills.is_empty() {
15279            return 0;
15280        }
15281        // Skills should share a phi_layer; warn (not fail) if they don't.
15282        let phi = skills[0].phi_layer;
15283        if skills.iter().any(|s| s.phi_layer != phi) {
15284            tracing::warn!("routable skills disagree on phi_layer; using {phi}");
15285        }
15286        let n = skills.len();
15287        self.set_dyn_phi_layer(Some(phi));
15288        self.dyn_router = Some(DynRouter::new(skills));
15289        n
15290    }
15291
15292    /// Human-readable switch log from the last dynamic-routed generation.
15293    pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
15294        self.dyn_router
15295            .as_ref()
15296            .map(|r| r.switches.clone())
15297            .unwrap_or_default()
15298    }
15299
15300    /// LM head: hidden → logits [vocab_size]. The dominant matvec of
15301    /// every decode step — row-parallel on the worker pool.
15302    fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
15303        let _mimo_q8 = self.mimo_moe.is_on()
15304            .then(crate::qtensor::enter_full_gpu_q8_scope);
15305        let rows = self.weights.lm_head.rows();
15306        let mut logits = attention::take_buf(rows.min(self.vocab_size));
15307        // Banked MiMo uses the same exact projection family for the
15308        // plain/draft head and the batched verification head. Read both
15309        // scale planes in-place instead of preparing per-op scale buffers.
15310        let served = self.mimo_moe.is_on() && crate::gpu::mimo_q8_short_enabled()
15311            && rows == self.vocab_size && !self.weights.lm_head.has_prism_contract()
15312            && self.weights.lm_head.graph_weight().is_some_and(|(model, idx, kind, _)| {
15313                kind == 7 && crate::gpu::q82_short_rows(model, idx, hidden, 1,
15314                    rows, self.hidden_size, &mut logits)
15315            });
15316        if !served {
15317            self.weights.lm_head.matvec(hidden, &mut logits, self.pool.as_deref());
15318        }
15319        logits.resize(self.vocab_size, 0.0);
15320        if let Some(m) = self.logit_multiplier {
15321            for l in logits.iter_mut() {
15322                *l *= m;
15323            }
15324        }
15325        if let Some(c) = self.final_softcap {
15326            for l in logits.iter_mut() {
15327                *l = c * (*l / c).tanh();
15328            }
15329        }
15330        if let Some(cm) = self.head_clusters.as_ref() {
15331            self.hierarchical_head_logprobs(hidden, cm, &mut logits);
15332        }
15333        logits
15334    }
15335
15336    /// Two-level head (Cortiq Embryo): in place, logits[v] ← log p(v) =
15337    /// (lc[c] − lse(lc)) + (logit[v] − lse over v's cluster block), c = v / S.
15338    fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
15339        let h = hidden.len();
15340        let ncl = cm.len() / h.max(1);
15341        if ncl == 0 || logits.len() % ncl != 0 {
15342            return;
15343        }
15344        let cs = logits.len() / ncl;
15345        // cluster logits + log-softmax
15346        let mut lc = vec![0.0f32; ncl];
15347        for c in 0..ncl {
15348            let row = &cm[c * h..(c + 1) * h];
15349            let mut s = 0.0f32;
15350            for j in 0..h {
15351                s += row[j] * hidden[j];
15352            }
15353            lc[c] = s;
15354        }
15355        let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15356        let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
15357        for c in 0..ncl {
15358            let blk = &mut logits[c * cs..(c + 1) * cs];
15359            let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15360            let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
15361            let add = lc[c] - lse - bl;
15362            for v in blk.iter_mut() {
15363                *v += add;
15364            }
15365        }
15366    }
15367
15368    /// Prefill `ids` and return the next-token logits — what the model
15369    /// would predict next, WITHOUT committing to generation (introspection
15370    /// for `cortiq explain`). Clears and repopulates the KV cache; leaves
15371    /// the active overlay untouched.
15372    pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
15373        #[cfg(target_os = "macos")]
15374        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
15375        self.clear_sequence_state();
15376        // This helper is used by the pooled classification endpoint, where
15377        // every request is a fresh sequence. The shared reset also clears the
15378        // wgpu token graph's device-side recurrent state.
15379        crate::gpu::graph_race_begin_generation();
15380        if task_mask.is_none() {
15381            self.o1_begin();
15382        }
15383        let mut hidden = vec![0.0f32; self.hidden_size];
15384        for (pos, &id) in ids.iter().enumerate() {
15385            let emb = self.embed_single(id);
15386            hidden = self.forward_layers(&emb, pos, task_mask);
15387        }
15388        if let Err(err) = self.o1_seal_checked() {
15389            self.o1_fail(err);
15390        }
15391        // Stacks that own their head (V4, V4.1, Qwen3.8-Flash-Next, GLM-5)
15392        // return a zero hidden and hand the logits out of band.
15393        if let Some(logits) = self.graph_logits.take() {
15394            return logits;
15395        }
15396        inference::rms_norm_into(
15397            &hidden,
15398            &self.weights.final_norm,
15399            self.rms_eps,
15400            self.norm_style,
15401            &mut self.ws.n1,
15402        );
15403        self.lm_head_forward(&self.ws.n1)
15404    }
15405}
15406
15407/// Convenience: deterministic tiny pipeline for tests.
15408pub fn create_test_pipeline(
15409    hidden_size: usize,
15410    intermediate_size: usize,
15411    num_heads: usize,
15412    num_kv_heads: usize,
15413    head_dim: usize,
15414    num_layers: usize,
15415    vocab_size: usize,
15416) -> Pipeline {
15417    // Small pseudo-random weights: constant weights make attention
15418    // degenerate and hide indexing bugs.
15419    let synth = |n: usize, salt: usize| -> Vec<f32> {
15420        (0..n)
15421            .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
15422            .collect()
15423    };
15424    let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
15425        QTensor::from_f32(synth(rows * cols, salt), rows, cols)
15426    };
15427    let layer_weights: Vec<LayerWeights> = (0..num_layers)
15428        .map(|li| LayerWeights {
15429            input_norm: vec![1.0; hidden_size],
15430            post_norm: vec![1.0; hidden_size],
15431            attn_out_norm: None,
15432            ffn_out_norm: None,
15433            layer_scale: None,
15434            ffn: FfnKind::Dense(DenseFfn {
15435                gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
15436                up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
15437                down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
15438                act: Act::Silu,
15439                down_t: None,
15440                segs: Vec::new(),
15441            }),
15442            attn: AttnKind::Full {
15443                bias: None,
15444                wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
15445                wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
15446                wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
15447                wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
15448                q_norm: None,
15449                k_norm: None,
15450                output_gate: false,
15451                softplus_gate: None,
15452            },
15453        })
15454        .collect();
15455
15456    Pipeline::new(
15457        Tokenizer::byte_level(),
15458        PipelineWeights {
15459            embed_tokens: qt(vocab_size, hidden_size, 100),
15460            layers: layer_weights,
15461            lm_head: qt(vocab_size, hidden_size, 200),
15462            final_norm: vec![1.0; hidden_size],
15463        },
15464        hidden_size,
15465        intermediate_size,
15466        num_heads,
15467        num_kv_heads,
15468        head_dim,
15469        num_layers,
15470        num_layers, // physical_layers = num_layers (non-looped)
15471        false,      // loop_final_norm
15472        vocab_size,
15473        1e-6,
15474        10_000.0,
15475        NormStyle::Qwen,
15476        4096,
15477        SamplerConfig {
15478            seed: Some(42),
15479            ..Default::default()
15480        },
15481    )
15482}
15483
15484/// Batched dense-FFN: gate/up/down via matmat (element-wise the same
15485/// math as b × dense_ffn — the same dot kernels).
15486/// One mask bit, LSB-first per byte — `TaskMask::ffn_active_indices`'s
15487/// convention.
15488#[inline]
15489fn mask_bit(row: &[u8], j: usize) -> bool {
15490    (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
15491}
15492
15493/// Zero the CLOSED neurons' activations in a [rows × inter] panel — the
15494/// masked-inference fast path's whole trick: full fused quant compute,
15495/// then the mask lands on the ACTIVATIONS, which is arithmetically the
15496/// pruned network without touching a quantized weight byte. Whole open
15497/// bytes (0xFF = 8 open neurons) skip in one test.
15498/// `CMF_FFN_MASK_GAIN` — Patent 12 FIG. 4, variance-preserving
15499/// rescaling: truncation removes a share of the layer's output energy,
15500/// so the survivors are scaled up to put the variance back where the
15501/// downstream norm expects it. A scalar here; per layer it is
15502/// `sqrt(total energy / kept energy)`.
15503fn mask_gain() -> f32 {
15504    static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
15505    *G.get_or_init(|| {
15506        std::env::var("CMF_FFN_MASK_GAIN")
15507            .ok()
15508            .and_then(|v| v.parse().ok())
15509            .unwrap_or(1.0)
15510    })
15511}
15512
15513fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
15514    // With CMF_FFN_MEANFILL a closed neuron contributes its average
15515    // instead of nothing — same bytes read, one constant restored.
15516    let fill = meanfill().and_then(|(i, v)| {
15517        let li = crate::gpu::cur_layer();
15518        (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
15519    });
15520    for r in 0..rows {
15521        let base = r * inter;
15522        for (bi, &byte) in row.iter().enumerate() {
15523            if byte == 0xFF {
15524                continue;
15525            }
15526            let j0 = bi * 8;
15527            for bit in 0..8 {
15528                let j = j0 + bit;
15529                if j < inter && byte & (1 << bit) == 0 {
15530                    g[base + j] = fill.map_or(0.0, |f| f[j]);
15531                }
15532            }
15533        }
15534    }
15535    let gain = mask_gain();
15536    if gain != 1.0 {
15537        for v in g[..rows * inter].iter_mut() {
15538            *v *= gain;
15539        }
15540    }
15541}
15542
15543/// True when neuron `i`'s bit is set (no mask = everything runs).
15544#[inline]
15545fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
15546    row.is_none_or(|r| mask_bit(r, i))
15547}
15548
15549/// Every bit below `n` set — the common case for a tube file's CORE,
15550/// where only the tube bits vary per task.
15551fn all_bits_on(row: &[u8], n: usize) -> bool {
15552    (0..n).all(|i| mask_bit(row, i))
15553}
15554
15555/// `CMF_TUBE_TOPK` — how many tubes a TOKEN may open (0 = the task mask
15556/// decides alone). This is the dense FFN read as a mixture: the tubes
15557/// are the experts a k-means over `gate_proj` rows found, and the token
15558/// picks among them. `CMF_TUBE_SCORE=gate` scores a tube by its own
15559/// gate (realizable: only `up`/`down` of the losers go unread),
15560/// `=oracle` scores by the true `silu(gate)·up` mass (the ceiling —
15561/// only `down` is saved, and the selection has read what it predicts).
15562fn tube_topk() -> usize {
15563    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
15564    *K.get_or_init(|| {
15565        std::env::var("CMF_TUBE_TOPK")
15566            .ok()
15567            .and_then(|v| v.parse().ok())
15568            .unwrap_or(0)
15569    })
15570}
15571
15572fn tube_score_oracle() -> bool {
15573    static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15574    *O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
15575}
15576
15577/// The routed arm of `tube_ffn`: a token opens only its best `k` tubes.
15578/// At `b == 1` (decode) the losers are genuinely never read — that is
15579/// the speed. At `b > 1` (the scoring sweep) every tube is computed and
15580/// the losers' activations are zeroed instead: same arithmetic, so the
15581/// perplexity is the routed model's, measured without a per-token
15582/// gather in the middle of a GEMM.
15583fn tube_ffn_routed(
15584    d: &DenseFfn,
15585    xs: &[f32],
15586    b: usize,
15587    pool: Option<&Pool>,
15588    mask_row: Option<&[u8]>,
15589    k: usize,
15590) -> Vec<f32> {
15591    let hidden = d.down_proj.rows();
15592    let core = d.gate_proj.rows();
15593    let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
15594    let mut out = match (b, core_full, mask_row) {
15595        (1, true, _) => dense_ffn(d, xs, pool),
15596        (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
15597        (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
15598        (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
15599    };
15600    let cand: Vec<usize> = (0..d.segs.len())
15601        .filter(|&i| tube_bit(mask_row, d.segs[i].start))
15602        .collect();
15603    if cand.is_empty() {
15604        return out;
15605    }
15606    // gate (and, where the score or the batch needs it, up) per tube.
15607    // The SCORE is taken at the point the serving path could take it:
15608    // off the gate alone, or off the finished activation for the oracle.
15609    let oracle = tube_score_oracle();
15610    let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
15611    let mut scores = vec![0f32; b * cand.len()];
15612    for (ci, &i) in cand.iter().enumerate() {
15613        let seg = &d.segs[i];
15614        let w = seg.width;
15615        let mut g = vec![0.0f32; b * w];
15616        if b == 1 {
15617            seg.gate.matvec(xs, &mut g, pool);
15618        } else {
15619            seg.gate.matmat(xs, b, &mut g, pool);
15620        }
15621        for v in g.iter_mut() {
15622            *v = Act::Silu.combine(*v, 1.0);
15623        }
15624        if !oracle {
15625            for t in 0..b {
15626                scores[t * cand.len() + ci] =
15627                    g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15628            }
15629        }
15630        if oracle || b > 1 {
15631            let mut u = vec![0.0f32; b * w];
15632            if b == 1 {
15633                seg.up.matvec(xs, &mut u, pool);
15634            } else {
15635                seg.up.matmat(xs, b, &mut u, pool);
15636            }
15637            for (a, &v) in g.iter_mut().zip(u.iter()) {
15638                *a *= v;
15639            }
15640            if oracle {
15641                for t in 0..b {
15642                    scores[t * cand.len() + ci] =
15643                        g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15644                }
15645            }
15646        }
15647        acts.push(g);
15648    }
15649    // per-token scores and the winners
15650    let keep = k.min(cand.len());
15651    let mut scratch: Vec<f32> = Vec::new();
15652    for t in 0..b {
15653        let mut sc: Vec<(f32, usize)> = (0..cand.len())
15654            .map(|ci| (scores[t * cand.len() + ci], ci))
15655            .collect();
15656        sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
15657        let mut alive = vec![false; cand.len()];
15658        for &(_, ci) in sc.iter().take(keep) {
15659            alive[ci] = true;
15660        }
15661        if b > 1 {
15662            for (ci, a) in acts.iter_mut().enumerate() {
15663                if !alive[ci] {
15664                    let w = d.segs[cand[ci]].width;
15665                    a[t * w..(t + 1) * w].fill(0.0);
15666                }
15667            }
15668        } else {
15669            // decode: finish only the winners — the losers' up/down
15670            // (and, with the gate score, everything but their gate)
15671            // are never touched.
15672            for (ci, &i) in cand.iter().enumerate() {
15673                if !alive[ci] {
15674                    continue;
15675                }
15676                let seg = &d.segs[i];
15677                let w = seg.width;
15678                let g = &mut acts[ci];
15679                if !tube_score_oracle() {
15680                    scratch.clear();
15681                    scratch.resize(w, 0.0);
15682                    seg.up.matvec(xs, &mut scratch, pool);
15683                    for (a, &v) in g.iter_mut().zip(scratch.iter()) {
15684                        *a *= v;
15685                    }
15686                }
15687                let mut acc = vec![0.0f32; hidden];
15688                seg.down.matvec(g, &mut acc, pool);
15689                for (o, a) in out.iter_mut().zip(&acc) {
15690                    *o += *a;
15691                }
15692            }
15693        }
15694    }
15695    if b > 1 {
15696        for (ci, &i) in cand.iter().enumerate() {
15697            let seg = &d.segs[i];
15698            let mut acc = vec![0.0f32; b * hidden];
15699            seg.down.matmat(&acts[ci], b, &mut acc, pool);
15700            for (o, a) in out.iter_mut().zip(&acc) {
15701                *o += *a;
15702            }
15703        }
15704    }
15705    out
15706}
15707
15708/// FFN of a defragged tube layer: the always-on core plus the tubes the
15709/// task mask switches on. Each tube is a normal tensor triple, so the
15710/// same kernels run it and an inactive tube's bytes are never read —
15711/// that is the whole point of the defrag (a scattered mask cannot skip
15712/// bytes; a contiguous one is just a smaller matrix).
15713fn tube_ffn(
15714    d: &DenseFfn,
15715    xs: &[f32],
15716    b: usize,
15717    pool: Option<&Pool>,
15718    mask_row: Option<&[u8]>,
15719) -> Vec<f32> {
15720    if tube_topk() > 0 {
15721        return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
15722    }
15723    let hidden = d.down_proj.rows();
15724    let core = d.gate_proj.rows();
15725    let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
15726    let mut out = match (b, core_full, mask_row) {
15727        (1, true, _) => dense_ffn(d, xs, pool),
15728        (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
15729        (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
15730        (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
15731    };
15732    TUBE_SCRATCH.with(|sc| {
15733        let mut sc = sc.borrow_mut();
15734        let [g, u, acc] = &mut *sc;
15735        for seg in &d.segs {
15736            if !tube_bit(mask_row, seg.start) {
15737                continue;
15738            }
15739            let w = seg.width;
15740            g.resize(b * w, 0.0);
15741            if b == 1
15742                && d.act == Act::Silu
15743                && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
15744            {
15745                // g holds silu(gate)·up.
15746            } else {
15747                u.resize(b * w, 0.0);
15748                if b == 1 {
15749                    QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
15750                } else {
15751                    seg.gate.matmat(xs, b, g, pool);
15752                    seg.up.matmat(xs, b, u, pool);
15753                }
15754                for i in 0..b * w {
15755                    g[i] = d.act.combine(g[i], u[i]);
15756                }
15757            }
15758            acc.resize(b * hidden, 0.0);
15759            acc.fill(0.0);
15760            if b == 1 {
15761                seg.down.matvec(g, acc, pool);
15762            } else {
15763                seg.down.matmat(g, b, acc, pool);
15764            }
15765            for (o, a) in out.iter_mut().zip(acc.iter()) {
15766                *o += *a;
15767            }
15768        }
15769        out
15770    })
15771}
15772
15773thread_local! {
15774    /// gate / up / down-accumulator scratch for the tube loop — a tube
15775    /// runs once per layer per token, and a fresh Vec each time is a
15776    /// malloc per tube per layer per token.
15777    static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
15778        const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
15779}
15780
15781fn dense_ffn_batch(
15782    d: &DenseFfn,
15783    xs: &[f32],
15784    b: usize,
15785    pool: Option<&Pool>,
15786    mask_row: Option<&[u8]>,
15787) -> Vec<f32> {
15788    let inter = d.gate_proj.rows();
15789    let hidden = d.down_proj.rows();
15790    // Fused on-device SwiGLU when the device is in play: three separate
15791    // `matmat` calls are three round trips per layer, and the gate/up
15792    // panels (b × inter — 22 MB each at a 512-token chunk) cross the bus
15793    // twice for nothing. The kernel already existed for the image DiT;
15794    // the LLM prefill was simply never wired to it. A task mask needs the
15795    // activations on the host between the halves, so it keeps the CPU
15796    // arm below.
15797    if mask_row.is_none()
15798        && d.act == Act::Silu
15799        && b >= 32
15800        && crate::gpu::enabled_here()
15801        && !crate::gpu::mm_killed()
15802        // The refit pass needs this layer's activations on the host; the
15803        // fused chain keeps them on the device. Refusing it here costs
15804        // one round trip and keeps every GEMM on the card — the
15805        // alternative was running the whole calibration on the CPU.
15806        && refit_dir().is_none()
15807        // Same for the mass/hit probes. The accumulator at the bottom of
15808        // this function only sees `g` when `g` came back to the host, so
15809        // a fused batch would leave it summing nothing — a probe that
15810        // reports zeros rather than failing, which is worse.
15811        && !ffn_probe_active()
15812    {
15813        if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
15814            d.gate_proj.mapped_q4t(),
15815            d.up_proj.mapped_q4t(),
15816            d.down_proj.mapped_q4t(),
15817        ) {
15818            let mut out = vec![0.0f32; b * hidden];
15819            if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
15820                return out;
15821            }
15822        }
15823        // The q4tp twin (same kernel family, scale from the row ladder) —
15824        // the DiT has run it in production since the pipeline containers;
15825        // the LLM prefill was simply never wired to it, so a q4tp model's
15826        // prefill panels stayed on the CPU.
15827        if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
15828            d.gate_proj.mapped_q4tp(),
15829            d.up_proj.mapped_q4tp(),
15830            d.down_proj.mapped_q4tp(),
15831        ) {
15832            let mut out = vec![0.0f32; b * hidden];
15833            if crate::gpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
15834                return out;
15835            }
15836        }
15837    }
15838    let mut g = vec![0.0f32; b * inter];
15839    d.gate_proj.matmat(xs, b, &mut g, pool);
15840    let mut u = vec![0.0f32; b * inter];
15841    d.up_proj.matmat(xs, b, &mut u, pool);
15842    if gate_topk() > 0 && d.act == Act::Silu {
15843        for t in 0..b {
15844            let row = &mut g[t * inter..(t + 1) * inter];
15845            for v in row.iter_mut() {
15846                *v = Act::Silu.combine(*v, 1.0);
15847            }
15848            keep_top_k(row, gate_topk());
15849        }
15850        for i in 0..b * inter {
15851            g[i] *= u[i];
15852        }
15853    } else {
15854        for i in 0..b * inter {
15855            g[i] = d.act.combine(g[i], u[i]);
15856        }
15857    }
15858    if let Some(row) = mask_row {
15859        zero_masked_cols(&mut g, b, inter, row);
15860    }
15861    if oracle_topk() > 0 {
15862        for t in 0..b {
15863            keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
15864        }
15865    }
15866    let mut out = vec![0.0f32; b * hidden];
15867    d.down_proj.matmat(&g, b, &mut out, pool);
15868    if refit_dir().is_some() {
15869        let li = crate::gpu::cur_layer();
15870        if li >= 0 {
15871            refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
15872        }
15873    }
15874    // The DTG-MA probe, on the batched path: one prefill sweep gives the
15875    // same per-neuron statistic the per-position probe does, and on a 27B
15876    // that is minutes instead of hours.
15877    FFN_PROBE.with(|pr| {
15878        if let Some(acc) = pr.borrow_mut().as_mut() {
15879            let li = crate::gpu::cur_layer();
15880            if li < 0 {
15881                return;
15882            }
15883            let Some(row) = acc.get_mut(li as usize) else {
15884                return;
15885            };
15886            let sq = probe_sq();
15887            for t in 0..b {
15888                for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
15889                    *a += if sq {
15890                        (v as f64) * (v as f64)
15891                    } else {
15892                        (v as f64).abs()
15893                    };
15894                }
15895            }
15896        }
15897    });
15898    out
15899}
15900
15901/// Batched MoE-FFN: router batched, positions are GROUPED by expert —
15902/// an expert's weights are read once for all its positions in the chunk
15903/// (the main prefill-GEMM win on MoE: 960MB/token of 35B experts).
15904/// Accumulate per-channel activation energy for `CMF_RMS_TRACE`.
15905fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
15906    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15907    static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15908    let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
15909    let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
15910    if (!on && !dump) || b == 0 {
15911        return;
15912    }
15913    let hidden = xs.len() / b;
15914    if on {
15915        let mut acc = m.act_sq.borrow_mut();
15916        if acc.len() < hidden {
15917            acc.resize(hidden, 0.0);
15918        }
15919        for t in 0..b {
15920            let row = &xs[t * hidden..(t + 1) * hidden];
15921            for (a, &v) in acc.iter_mut().zip(row) {
15922                *a += (v as f64) * (v as f64);
15923            }
15924        }
15925    }
15926    if dump {
15927        // Cap the capture: the covariance needs a few thousand rows, and a
15928        // whole prefill of every layer would be gigabytes for no extra rank.
15929        let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
15930            .ok()
15931            .and_then(|v| v.parse().ok())
15932            .unwrap_or(4096);
15933        let mut rows = m.act_rows.borrow_mut();
15934        if rows.len() < cap * hidden {
15935            let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
15936            rows.extend_from_slice(&xs[..take * hidden]);
15937        }
15938    }
15939}
15940
15941/// Send-able cursor over a Vec-of-Vecs: each pool worker writes only its
15942/// own slots (disjoint by construction in the caller).
15943#[derive(Clone, Copy)]
15944struct SendVecs(*mut Vec<f32>);
15945unsafe impl Send for SendVecs {}
15946unsafe impl Sync for SendVecs {}
15947impl SendVecs {
15948    #[inline]
15949    fn at(self, i: usize) -> *mut Vec<f32> {
15950        unsafe { self.0.add(i) }
15951    }
15952}
15953
15954fn moe_ffn_batch(
15955    m: &MoeFfn,
15956    xs: &[f32],
15957    b: usize,
15958    hidden: usize,
15959    pool: Option<&Pool>,
15960    allowed: Option<&[bool]>,
15961) -> Vec<f32> {
15962    accumulate_act(m, xs, b);
15963    let ne = m.experts.len();
15964    let mut logits = vec![0.0f32; b * ne];
15965    match &m.resonance {
15966        Some(r) => {
15967            let hdim = xs.len() / b.max(1);
15968            for bi in 0..b {
15969                r.scores(
15970                    &xs[bi * hdim..(bi + 1) * hdim],
15971                    &mut logits[bi * ne..(bi + 1) * ne],
15972                );
15973            }
15974        }
15975        None => m.router.matmat(xs, b, &mut logits, pool),
15976    }
15977
15978    // Assignments: expert → [(position, weight)] — same routing as
15979    // moe_ffn, per position (see `moe_route`).
15980    let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
15981    {
15982        let mut st = m.stats.borrow_mut();
15983        if st.len() < ne {
15984            st.resize(ne, 0);
15985        }
15986        for bi in 0..b {
15987            let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
15988            for &e in &idx {
15989                st[e] += 1;
15990                assign[e].push((bi, p[e] / wsum));
15991            }
15992        }
15993    }
15994
15995    let mut out = vec![0.0f32; b * hidden];
15996    let cols = m.experts[0].gate_proj.cols();
15997    let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
15998        let sb = list.len();
15999        let mut sub = vec![0.0f32; sb * cols];
16000        for (k, &(bi, _)) in list.iter().enumerate() {
16001            sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16002        }
16003        let eo = dense_ffn_batch(d, &sub, sb, pool, None);
16004        for (k, &(bi, w)) in list.iter().enumerate() {
16005            for i in 0..hidden {
16006                out[bi * hidden + i] += w * eo[k * hidden + i];
16007            }
16008        }
16009    };
16010    // Routed experts: the panels are TINY (b·top_k spread over every
16011    // expert — a few positions each), so a pool dispatch per expert is
16012    // pure barrier cost. Invert the parallelism: workers take WHOLE
16013    // experts (serial math inside), then one deterministic scatter in
16014    // expert order — the exact accumulation order the serial loop had.
16015    let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
16016    if pool.is_some() && active.len() >= 8 {
16017        let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
16018        {
16019            let panel_ptr = SendVecs(panels.as_mut_ptr());
16020            // Capture only the expert table: `m` itself carries RefCell
16021            // stats and must not cross the pool boundary.
16022            let experts = &m.experts;
16023            let (active_r, assign_r) = (&active, &assign);
16024            let inherit_cpu = crate::gpu::inherit_cpu_scope();
16025            let run = |start: usize, end: usize| {
16026                let _cpu_scope = inherit_cpu();
16027                for ai in start..end {
16028                    let e = active_r[ai];
16029                    let list = &assign_r[e];
16030                    let sb = list.len();
16031                    let mut sub = vec![0.0f32; sb * cols];
16032                    for (k, &(bi, _)) in list.iter().enumerate() {
16033                        sub[k * cols..(k + 1) * cols]
16034                            .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16035                    }
16036                    // SAFETY: each worker owns a disjoint panels[ai].
16037                    unsafe {
16038                        *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
16039                    }
16040                }
16041            };
16042            match pool {
16043                Some(p) => p.run_rows(active.len(), &run),
16044                None => run(0, active.len()),
16045            }
16046        }
16047        for (ai, &e) in active.iter().enumerate() {
16048            for (k, &(bi, w)) in assign[e].iter().enumerate() {
16049                let eo = &panels[ai][k * hidden..(k + 1) * hidden];
16050                for i in 0..hidden {
16051                    out[bi * hidden + i] += w * eo[i];
16052                }
16053            }
16054        }
16055    } else {
16056        for &e in &active {
16057            run_expert(&m.experts[e], &assign[e], &mut out);
16058        }
16059    }
16060    if let Some((se, gate)) = &m.shared {
16061        let all: Vec<(usize, f32)> = if let Some(gate) = gate {
16062            let mut gl = vec![0.0f32; b];
16063            gate.matmat(xs, b, &mut gl, pool);
16064            (0..b)
16065                .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
16066                .collect()
16067        } else {
16068            (0..b).map(|bi| (bi, 1.0)).collect()
16069        };
16070        run_expert(se, &all, &mut out);
16071    }
16072    out
16073}
16074
16075/// Decode-exact multi-token MoE — the MiMo speculative verify's FFN. Row
16076/// `r` of the result is bit-identical to `moe_ffn(m, x_r)` on the CPU
16077/// (`moe_ffn_cpu` → `moe_ffn_cpu_batched`): router matvec per row, the same
16078/// routing, the same int8 gate/up/SiLU and down terms
16079/// (`QTensor::moe_gate_up_rows` / `moe_down_rows`), and the row's experts
16080/// summed in ITS route order from 0. What the rows share is the weight
16081/// traffic: each routed expert is read once for every row that picked it.
16082/// (`moe_ffn_batch`, the prompt path, groups the same way but sums in
16083/// expert-index order and runs blocked kernels on wide groups — close, not
16084/// bit-equal to decode.) Any layer the kernels do not cover, or a device
16085/// that could answer `moe_ffn` itself, walks `moe_ffn` row by row.
16086fn moe_ffn_rows_exact(
16087    m: &MoeFfn,
16088    xs: &[f32],
16089    b: usize,
16090    hidden: usize,
16091    pool: Option<&Pool>,
16092) -> Vec<f32> {
16093    let mut out = vec![0.0f32; b * hidden];
16094    let per_row = |out: &mut [f32]| {
16095        for r in 0..b {
16096            let o = moe_ffn(m, &xs[r * hidden..(r + 1) * hidden], pool, None);
16097            out[r * hidden..(r + 1) * hidden].copy_from_slice(&o);
16098        }
16099    };
16100    let covered = !crate::gpu::enabled_here()
16101        && moe_batch_enabled()
16102        && m.shared.is_none()
16103        && m.resonance.is_none()
16104        && FFN_PROBE.with(|pr| pr.borrow().is_none())
16105        && m.experts.iter().all(|d| d.act == Act::Silu);
16106    if !covered {
16107        per_row(&mut out);
16108        return out;
16109    }
16110    let ne = m.experts.len();
16111    // Routing, row by row, exactly as `moe_ffn`.
16112    let mut routes: Vec<(Vec<usize>, Vec<f32>)> = Vec::with_capacity(b);
16113    for r in 0..b {
16114        let x = &xs[r * hidden..(r + 1) * hidden];
16115        accumulate_act(m, x, 1);
16116        let mut logits = vec![0.0f32; ne];
16117        m.router.matvec(x, &mut logits, pool);
16118        let (idx, p, wsum) = moe_route(&logits, m, None);
16119        {
16120            let mut st = m.stats.borrow_mut();
16121            if st.len() < ne {
16122                st.resize(ne, 0);
16123            }
16124            for &e in &idx {
16125                st[e] += 1;
16126            }
16127        }
16128        let w: Vec<f32> = idx
16129            .iter()
16130            .map(|&e| p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]))
16131            .collect();
16132        routes.push((idx, w));
16133    }
16134    if routes.iter().any(|(idx, _)| idx.is_empty()) {
16135        per_row(&mut out);
16136        return out;
16137    }
16138    // Group the (row, expert) picks by expert, in first-seen order.
16139    let mut experts: Vec<usize> = Vec::new();
16140    let mut groups: Vec<Vec<usize>> = Vec::new();
16141    for (r, (idx, _)) in routes.iter().enumerate() {
16142        for &e in idx {
16143            match experts.iter().position(|&x| x == e) {
16144                Some(g) => groups[g].push(r),
16145                None => {
16146                    experts.push(e);
16147                    groups.push(vec![r]);
16148                }
16149            }
16150        }
16151    }
16152    let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
16153    let inter = m.experts[experts[0]].gate_proj.rows();
16154    let pairs: Vec<(&QTensor, &QTensor)> = experts
16155        .iter()
16156        .map(|&e| (&m.experts[e].gate_proj, &m.experts[e].up_proj))
16157        .collect();
16158    let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
16159    if !QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut gs, pool) {
16160        per_row(&mut out);
16161        return out;
16162    }
16163    let downs: Vec<&QTensor> = experts.iter().map(|&e| &m.experts[e].down_proj).collect();
16164    let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
16165    let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; hidden]).collect();
16166    if !QTensor::moe_down_rows(&downs, &lens, &gs, &mut ds, pool) {
16167        per_row(&mut out);
16168        return out;
16169    }
16170    // Where each (row, expert) term landed in the flat pair list.
16171    let mut slot = std::collections::HashMap::with_capacity(n_pairs);
16172    let mut p = 0usize;
16173    for (g, &e) in experts.iter().enumerate() {
16174        for &r in &groups[g] {
16175            slot.insert((r, e), p);
16176            p += 1;
16177        }
16178    }
16179    for (r, (idx, w)) in routes.iter().enumerate() {
16180        let terms: Vec<(&[f32], f32)> = idx
16181            .iter()
16182            .zip(w)
16183            .map(|(&e, &we)| (ds[slot[&(r, e)]].as_slice(), we))
16184            .collect();
16185        let row = &mut out[r * hidden..(r + 1) * hidden];
16186        for (i, dst) in row.iter_mut().enumerate() {
16187            // `moe_down_many`'s per-row sum: from 0, in route order.
16188            let mut acc = 0f32;
16189            for (d, we) in &terms {
16190                acc += we * d[i];
16191            }
16192            *dst = acc;
16193        }
16194    }
16195    out
16196}
16197
16198thread_local! {
16199    /// gate/up activation scratch for the dense FFN paths (single uses
16200    /// two slots, the fused pair all four) — these were fresh
16201    /// intermediate-size Vecs on every layer of every token.
16202    static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
16203        const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
16204}
16205
16206/// Dense SwiGLU FFN through QTensor matvecs (any storage).
16207fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16208    // Per-token sparsity, when the file was built for it: gate first,
16209    // then only the chosen neurons' up/down rows leave the mmap.
16210    if gate_topk() > 0
16211        && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
16212    {
16213        return out;
16214    }
16215    // Whole-FFN GPU submit (этап 4.2 increment): gate → silu·up → down
16216    // chained in ONE command buffer with the intermediate activations
16217    // resident on the device — 3 per-op polls become 1 per layer. The
16218    // moe_block backend already implements exactly this chain; a dense
16219    // FFN is one expert with weight 1. Runtime probe: the chain still
16220    // pays one submit+poll per layer — alternate it against the pure-CPU
16221    // FFN and keep whichever is faster on this machine.
16222    // q1 FFNs offload at any practical size: the q1 CPU kernel is
16223    // compute-bound, so the UMA threshold logic does not apply — the
16224    // probe measures and decides either way.
16225    // The fused GPU block has no descriptor-aware Prism path: it would either
16226    // consume an unrotated activation or decline after inspecting the mixed
16227    // q2tp/q4tp tensors.  Do not let that structural refusal enter the FFN
16228    // probe's CPU_ONLY scope; the ordinary body below dispatches each matrix
16229    // through QTensor::matvec, which owns the signed FWHT + affine q2tp route.
16230    let prism_body = d.gate_proj.has_prism_contract()
16231        || d.up_proj.has_prism_contract()
16232        || d.down_proj.has_prism_contract();
16233    if !prism_body
16234        && crate::gpu::enabled_here()
16235        && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
16236    {
16237        let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
16238            crate::gpu::ProbeArm::Gpu
16239        } else {
16240            crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
16241        };
16242        match arm {
16243            crate::gpu::ProbeArm::Gpu => {
16244                let t0 = std::time::Instant::now();
16245                if let Some(out) = dense_ffn_gpu(d, x, pool) {
16246                    crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
16247                    return out;
16248                }
16249                // Declined: no timing exists, so say so. Silence here is
16250                // what left `ffn` undecided for 9000 calls and cost a
16251                // failed device attempt on half of them.
16252                crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
16253            }
16254            crate::gpu::ProbeArm::CpuTimed => {
16255                let t0 = std::time::Instant::now();
16256                let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16257                crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
16258                return out;
16259            }
16260            crate::gpu::ProbeArm::Cpu => {
16261                return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16262            }
16263        }
16264    }
16265    dense_ffn_cpu(d, x, pool)
16266}
16267
16268/// The pure-CPU dense-FFN body (also the fallback of every GPU refusal).
16269fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16270    let inter = d.gate_proj.rows();
16271    FFN_SCRATCH.with(|s| {
16272        let mut s = s.borrow_mut();
16273        let [g, u, ..] = &mut *s;
16274        g.resize(inter, 0.0);
16275        // Fused gate+up+silu: one dispatch, no separate silu pass.
16276        // Falls back to matvec_many + silu loop for unsupported dtypes.
16277        if gate_topk() > 0 {
16278            // Gate first, select, and only then pay for `up`: the
16279            // measurement arm computes both and zeroes the losers, which
16280            // is the same arithmetic.
16281            u.resize(inter, 0.0);
16282            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16283            for i in 0..inter {
16284                g[i] = Act::Silu.combine(g[i], 1.0);
16285            }
16286            keep_top_k(g, gate_topk());
16287            for i in 0..inter {
16288                g[i] *= u[i];
16289            }
16290        } else if d.act == Act::Silu && {
16291            let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16292            QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
16293        } {
16294            // g now holds silu(gate)·up directly.
16295        } else {
16296            u.resize(inter, 0.0);
16297            // Multi-matrix job: gate+up under one pool dispatch.
16298            let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16299            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16300            for i in 0..inter {
16301                g[i] = d.act.combine(g[i], u[i]);
16302            }
16303        }
16304        // DTG-MA bake probe (Patent 2): accumulate this layer's
16305        // per-neuron activation mass while a probe pass is active.
16306        // `CMF_FFN_PROBE_TOPK=k` switches the statistic from mass to a
16307        // HIT COUNT — how many tokens rank the neuron in their own top
16308        // k. Mass asks "how loud is this neuron overall", the count
16309        // asks "how often does this task actually need it", and the two
16310        // rank neurons differently whenever a few tokens are loud.
16311        FFN_PROBE.with(|pr| {
16312            if let Some(acc) = pr.borrow_mut().as_mut() {
16313                let li = crate::gpu::cur_layer();
16314                if li >= 0 {
16315                    if let Some(row) = acc.get_mut(li as usize) {
16316                        match probe_topk() {
16317                            0 if probe_sq() => {
16318                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16319                                    *a += (v as f64) * (v as f64);
16320                                }
16321                            }
16322                            0 if probe_signed() => {
16323                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16324                                    *a += v as f64;
16325                                }
16326                            }
16327                            0 => {
16328                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16329                                    *a += (v as f64).abs();
16330                                }
16331                            }
16332                            k => {
16333                                let n = g.len();
16334                                let k = k.min(n);
16335                                let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16336                                let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16337                                    b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16338                                });
16339                                let thr = *kth;
16340                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16341                                    if v.abs() >= thr {
16342                                        *a += 1.0;
16343                                    }
16344                                }
16345                            }
16346                        }
16347                    }
16348                }
16349            }
16350        });
16351        if oracle_topk() > 0 {
16352            keep_top_k(g, oracle_topk());
16353        }
16354        {
16355            let li = crate::gpu::cur_layer();
16356            if li >= 0 {
16357                adump_row(li as usize, g);
16358            }
16359        }
16360        let mut out = attention::take_buf(d.down_proj.rows());
16361        let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnDown);
16362        d.down_proj.matvec(g, &mut out, pool);
16363        out
16364    })
16365}
16366
16367/// Online accumulators for the AWNP refit of a narrowed FFN.
16368///
16369/// The refit needs `Gss = A_SᵀA_S` and `YA = YᵀA_S` per layer, where `A_S`
16370/// are the calibration activations of the KEPT neurons and `Y` the full
16371/// FFN output. Both are small enough to hold; the thing that is not is
16372/// the activations they are built from — a 27B layer would dump a
16373/// gigabyte per thousand tokens. So they are accumulated as the
16374/// calibration runs and written once at the end.
16375///
16376/// `CMF_FFN_REFIT=<dir>` holds `support.<L>.u32` (a u32 count then the
16377/// kept indices) for every layer to accumulate; `CMF_FFN_REFIT_FROM/TO`
16378/// bound the layer span so the accumulators fit in RAM.
16379pub struct RefitAcc {
16380    pub support: Vec<u32>,
16381    pub gss: Vec<f32>,
16382    pub ya: Vec<f32>,
16383    pub hidden: usize,
16384    pub tokens: u64,
16385    /// Activations staged transposed ([ns, t] and [hidden, t]) until the
16386    /// batch is worth a GEMM. The product costs `ns²` to move and add
16387    /// REGARDLESS of how many tokens went into it, so folding 16 chunks
16388    /// into one call cuts that cost 16× — it was 15 TB of traffic per
16389    /// calibration pass at one call per 256 tokens.
16390    pub buf_g: Vec<f32>,
16391    pub buf_o: Vec<f32>,
16392    pub buf_t: usize,
16393}
16394
16395/// The product buffer is SHARED across layers — one 473 MB allocation,
16396/// not one per layer (that was 30 GB of nothing on a 64-layer model).
16397/// It lives under the same lock as the accumulators.
16398type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
16399
16400static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
16401    std::sync::OnceLock::new();
16402
16403/// Is an FFN probe accumulator installed on this thread? The fused GPU
16404/// FFN must decline while one is, or the probe silently measures zero.
16405fn ffn_probe_active() -> bool {
16406    FFN_PROBE.with(|p| p.borrow().is_some())
16407}
16408
16409fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
16410    REFIT
16411        .get_or_init(|| {
16412            std::env::var("CMF_FFN_REFIT").ok().map(|d| {
16413                (
16414                    d,
16415                    std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
16416                )
16417            })
16418        })
16419        .as_ref()
16420}
16421
16422/// Accumulate one prefill panel into the layer's refit statistics.
16423fn refit_accumulate(
16424    li: usize,
16425    g: &[f32],
16426    b: usize,
16427    inter: usize,
16428    out: &[f32],
16429    hidden: usize,
16430    pool: Option<&Pool>,
16431) {
16432    let Some((dir, map)) = refit_dir() else {
16433        return;
16434    };
16435    static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16436    let (from, to) = *SPAN.get_or_init(|| {
16437        let g = |k: &str, d: usize| {
16438            std::env::var(k)
16439                .ok()
16440                .and_then(|v| v.parse().ok())
16441                .unwrap_or(d)
16442        };
16443        (
16444            g("CMF_FFN_REFIT_FROM", 0),
16445            g("CMF_FFN_REFIT_TO", usize::MAX),
16446        )
16447    });
16448    if li < from || li > to {
16449        return;
16450    }
16451    let mut guard = map.lock().unwrap();
16452    let (map, shared) = &mut *guard;
16453    let acc = match map.entry(li) {
16454        std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
16455        std::collections::hash_map::Entry::Vacant(e) => {
16456            let path = format!("{dir}/support.{li}.u32");
16457            let Ok(bytes) = std::fs::read(&path) else {
16458                eprintln!("refit: no {path} — layer {li} skipped");
16459                return;
16460            };
16461            let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
16462            let support: Vec<u32> = bytes[4..4 + n * 4]
16463                .chunks_exact(4)
16464                .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16465                .collect();
16466            eprintln!(
16467                "refit: layer {li} support {n} ({:.0} MB of accumulator)",
16468                (n * n + hidden * n) as f64 * 4.0 / 1e6
16469            );
16470            e.insert(RefitAcc {
16471                gss: vec![0.0; n * n],
16472                ya: vec![0.0; hidden * n],
16473                buf_g: Vec::new(),
16474                buf_o: Vec::new(),
16475                buf_t: 0,
16476                support,
16477                hidden,
16478                tokens: 0,
16479            })
16480        }
16481    };
16482    let ns = acc.support.len();
16483    // Stage this chunk transposed; the GEMM fires once the batch is full.
16484    let cap = refit_batch();
16485    if acc.buf_g.is_empty() {
16486        acc.buf_g = vec![0.0; ns * cap];
16487        acc.buf_o = vec![0.0; hidden * cap];
16488    }
16489    let take = b.min(cap - acc.buf_t);
16490    for t in 0..take {
16491        let col = acc.buf_t + t;
16492        for (j, &n) in acc.support.iter().enumerate() {
16493            acc.buf_g[j * cap + col] = g[t * inter + n as usize];
16494        }
16495        for h in 0..hidden {
16496            acc.buf_o[h * cap + col] = out[t * hidden + h];
16497        }
16498    }
16499    acc.buf_t += take;
16500    acc.tokens += take as u64;
16501    if acc.buf_t < cap {
16502        return;
16503    }
16504    let bt = acc.buf_t;
16505    acc.buf_t = 0;
16506    // The GEMM WRITES its C (it zeroes the accumulators it uses), so the
16507    // chunk product lands in scratch and is added on — the one thing that
16508    // silently turns a Gram over 13 000 tokens into a Gram over 256.
16509    // Both products are `C[n, m] += X[n, b] · Yᵀ[b, m]` with X and Y
16510    // stored row-major [·, b] — exactly `gemm_nt_f32`'s shape, so the
16511    // card does them when it is up (this is the whole calibration's
16512    // cost: O(|S|²) per token, 2.9 PFLOP for a 27B pass). The tiled CPU
16513    // loop stays as the fallback. Neither accumulates, so the product
16514    // lands in scratch and is added on.
16515    let RefitAcc {
16516        gss,
16517        ya,
16518        buf_g,
16519        buf_o,
16520        ..
16521    } = acc;
16522    let need = (ns * ns).max(hidden * ns);
16523    if shared.len() < need {
16524        shared.resize(need, 0.0);
16525    }
16526    let scratch = &mut shared[..];
16527    let _ = bt;
16528    if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
16529        add_into(gss, &scratch[..ns * ns], pool);
16530        if crate::gpu::gemm_nt_f32_transient(
16531            buf_o,
16532            buf_g,
16533            &mut scratch[..hidden * ns],
16534            hidden,
16535            cap,
16536            ns,
16537        ) {
16538            add_into(ya, &scratch[..hidden * ns], pool);
16539        } else {
16540            accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16541        }
16542    } else {
16543        accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
16544        accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16545    }
16546    // No zeroing: the batch is always filled exactly (cap is a multiple
16547    // of the prefill chunk), and a memset of 178 MB a layer would cost
16548    // more than the GEMM.
16549}
16550
16551/// `CMF_FFN_REFIT_BATCH` — tokens staged before each GEMM (default 4096).
16552fn refit_batch() -> usize {
16553    static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16554    *B.get_or_init(|| {
16555        std::env::var("CMF_FFN_REFIT_BATCH")
16556            .ok()
16557            .and_then(|v| v.parse().ok())
16558            .unwrap_or(4096)
16559    })
16560}
16561
16562/// `c[m, n] += Σ_t left[m, t]·right[n, t]` — both operands transposed,
16563/// the CPU fallback for the staged batch.
16564fn accum_outer_t(
16565    c: &mut [f32],
16566    m: usize,
16567    n: usize,
16568    b: usize,
16569    left: &[f32],
16570    right: &[f32],
16571    pool: Option<&Pool>,
16572) {
16573    let ptr = SendMut(c.as_mut_ptr());
16574    let body = |i: usize| {
16575        let ptr = &ptr;
16576        let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
16577        for t in 0..b {
16578            let a = left[i * b + t];
16579            if a == 0.0 {
16580                continue;
16581            }
16582            for (j, o) in row.iter_mut().enumerate() {
16583                *o += a * right[j * b + t];
16584            }
16585        }
16586    };
16587    match pool {
16588        Some(p) if m > 1 => p.run_rows(m, &|s, e| {
16589            for i in s..e {
16590                body(i);
16591            }
16592        }),
16593        _ => {
16594            for i in 0..m {
16595                body(i);
16596            }
16597        }
16598    }
16599}
16600
16601/// `dst += src`, spread over the pool — at 118 M floats a layer this is
16602/// not a loop to leave on one core.
16603fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
16604    let n = dst.len().min(src.len());
16605    match pool {
16606        Some(p) if n >= 1 << 16 => {
16607            let ptr = SendMut(dst.as_mut_ptr());
16608            let f = |s: usize, e: usize| {
16609                let ptr = &ptr;
16610                for blk in s..e {
16611                    let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
16612                    for i in a..b {
16613                        unsafe { *ptr.0.add(i) += src[i] };
16614                    }
16615                }
16616            };
16617            p.run_rows(n.div_ceil(4096), &f);
16618        }
16619        _ => {
16620            for (d, v) in dst.iter_mut().zip(&src[..n]) {
16621                *d += *v;
16622            }
16623        }
16624    }
16625}
16626
16627/// `c[m, n] += Σ_t left[t, m]·right[t, n]`, with `left` stored [m, t] and
16628/// `right` [t, n]. Tiled over the rows of `c` so a tile stays in cache
16629/// while each token's `right` row streams past it once, and parallel
16630/// over tiles.
16631fn accum_outer(
16632    c: &mut [f32],
16633    m: usize,
16634    n: usize,
16635    b: usize,
16636    left: &[f32],
16637    right: &[f32],
16638    pool: Option<&Pool>,
16639) {
16640    const TILE: usize = 32;
16641    let tiles = m.div_ceil(TILE);
16642    let cp = SendMut(c.as_mut_ptr());
16643    let body = |ti: usize| {
16644        let cp = &cp;
16645        let i0 = ti * TILE;
16646        let i1 = (i0 + TILE).min(m);
16647        for t in 0..b {
16648            let r = &right[t * n..t * n + n];
16649            for i in i0..i1 {
16650                let a = left[i * b + t];
16651                if a == 0.0 {
16652                    continue;
16653                }
16654                // SAFETY: tiles partition c's rows; workers never overlap.
16655                let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
16656                for (o, v) in row.iter_mut().zip(r) {
16657                    *o += a * *v;
16658                }
16659            }
16660        }
16661    };
16662    match pool {
16663        Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
16664            for ti in s..e {
16665                body(ti);
16666            }
16667        }),
16668        _ => {
16669            for ti in 0..tiles {
16670                body(ti);
16671            }
16672        }
16673    }
16674}
16675
16676/// Write what the calibration accumulated: `gss.<L>.f32` and `ya.<L>.f32`.
16677pub fn refit_flush() -> usize {
16678    let Some((dir, map)) = refit_dir() else {
16679        return 0;
16680    };
16681    let guard = map.lock().unwrap();
16682    let mut n = 0;
16683    for (li, acc) in guard.0.iter() {
16684        // A silently truncated write here is a Gram that reshapes to
16685        // nothing an hour later — say it out loud instead.
16686        let w = |name: &str, v: &[f32]| {
16687            let path = format!("{dir}/{name}.{li}.f32");
16688            let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
16689            match std::fs::write(&path, &bytes) {
16690                Ok(()) => {}
16691                Err(e) => eprintln!(
16692                    "refit: FAILED to write {path} ({} MB): {e}",
16693                    bytes.len() / 1_000_000
16694                ),
16695            }
16696        };
16697        w("gss", &acc.gss);
16698        w("ya", &acc.ya);
16699        println!(
16700            "refit L{li}: {} support, {} tokens, hidden {}",
16701            acc.support.len(),
16702            acc.tokens,
16703            acc.hidden
16704        );
16705        n += 1;
16706    }
16707    n
16708}
16709
16710/// `CMF_FFN_ADUMP=<prefix>` — append every probed token's FFN activation
16711/// row to `<prefix>.<layer>.f16`. The co-activation record: which
16712/// neurons fire together, which is what a tube has to group if a token
16713/// is ever going to open one tube instead of sixteen.
16714fn adump_row(li: usize, g: &[f32]) {
16715    use std::io::Write as _;
16716    static FILES: std::sync::OnceLock<
16717        Option<(
16718            String,
16719            std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
16720        )>,
16721    > = std::sync::OnceLock::new();
16722    let Some((prefix, map)) = FILES
16723        .get_or_init(|| {
16724            std::env::var("CMF_FFN_ADUMP")
16725                .ok()
16726                .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
16727        })
16728        .as_ref()
16729    else {
16730        return;
16731    };
16732    // `CMF_FFN_ADUMP_FROM/_TO` narrow the dump to a layer span, so a big
16733    // calibration run fits on disk in a few passes instead of one.
16734    static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16735    let (from, to) = *SPAN.get_or_init(|| {
16736        let g = |k: &str, d: usize| {
16737            std::env::var(k)
16738                .ok()
16739                .and_then(|v| v.parse().ok())
16740                .unwrap_or(d)
16741        };
16742        (
16743            g("CMF_FFN_ADUMP_FROM", 0),
16744            g("CMF_FFN_ADUMP_TO", usize::MAX),
16745        )
16746    });
16747    if li < from || li > to {
16748        return;
16749    }
16750    let mut map = map.lock().unwrap();
16751    let f = map.entry(li).or_insert_with(|| {
16752        std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
16753    });
16754    let mut bytes = Vec::with_capacity(g.len() * 2);
16755    for v in g {
16756        bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
16757    }
16758    let _ = f.write_all(&bytes);
16759}
16760
16761/// `CMF_FFN_ORACLE_TOPK` — keep only the k largest |silu(g)·u| of each
16762/// token and zero the rest. Not a serving mode: it is the CEILING of
16763/// contextual sparsity — what a per-token router would be chasing —
16764/// measured by cheating, since the selection reads the very activations
16765/// it would have to predict.
16766fn oracle_topk() -> usize {
16767    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16768    *K.get_or_init(|| {
16769        std::env::var("CMF_FFN_ORACLE_TOPK")
16770            .ok()
16771            .and_then(|v| v.parse().ok())
16772            .unwrap_or(0)
16773    })
16774}
16775
16776/// `CMF_FFN_GATE_TOPK` — the REALIZABLE cousin of the oracle: rank the
16777/// neurons by their gate alone (which the kernel has computed anyway
16778/// before it reads `up`), keep the k best, and drop the rest. Every
16779/// dropped neuron's `up` row and `down` column stay unread, so this is
16780/// the sparsity a serving path can actually take without a router.
16781fn gate_topk() -> usize {
16782    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16783    *K.get_or_init(|| {
16784        std::env::var("CMF_FFN_GATE_TOPK")
16785            .ok()
16786            .and_then(|v| v.parse().ok())
16787            .unwrap_or(0)
16788    })
16789}
16790
16791/// `CMF_FFN_GATE_BLOCK` — select in blocks of B neurons instead of one
16792/// by one. A scattered per-neuron choice cannot be read efficiently (a
16793/// row at a time, no prefetch runway); a block of 32 is a contiguous
16794/// 32-row slab of `up` and of the transposed `down`, which the ordinary
16795/// kernels stream. The question the measurement answers is what the
16796/// block costs in quality.
16797fn gate_block() -> usize {
16798    static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16799    *B.get_or_init(|| {
16800        std::env::var("CMF_FFN_GATE_BLOCK")
16801            .ok()
16802            .and_then(|v| v.parse().ok())
16803            .unwrap_or(1)
16804    })
16805}
16806
16807/// Zero all but the `k` largest BLOCKS (by summed square) of a row.
16808fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
16809    let n = g.len();
16810    let nb = n.div_ceil(block);
16811    let kb = (keep_n.div_ceil(block)).clamp(1, nb);
16812    if kb >= nb {
16813        return;
16814    }
16815    let mut score: Vec<f32> = (0..nb)
16816        .map(|b| {
16817            g[b * block..((b + 1) * block).min(n)]
16818                .iter()
16819                .map(|v| v * v)
16820                .sum::<f32>()
16821        })
16822        .collect();
16823    let mut ord = score.clone();
16824    let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
16825        b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16826    });
16827    let thr = *kth;
16828    for b in 0..nb {
16829        if score[b] < thr {
16830            g[b * block..((b + 1) * block).min(n)].fill(0.0);
16831        }
16832    }
16833    score.clear();
16834}
16835
16836/// Zero all but the `k` largest magnitudes of one token's activation row.
16837fn keep_top_k(g: &mut [f32], k: usize) {
16838    if gate_block() > 1 {
16839        return keep_top_blocks(g, k, gate_block());
16840    }
16841    let n = g.len();
16842    if k == 0 || k >= n {
16843        return;
16844    }
16845    let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16846    let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16847        b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16848    });
16849    let thr = *kth;
16850    for v in g.iter_mut() {
16851        if v.abs() < thr {
16852            *v = 0.0;
16853        }
16854    }
16855}
16856
16857/// `CMF_FFN_PROBE_SQ` — accumulate Σa², so the dump divided by the token
16858/// count and square-rooted is the RMS activation trace Patent 12 weights
16859/// its matrices by.
16860fn probe_sq() -> bool {
16861    static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16862    *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
16863}
16864
16865/// `CMF_FFN_PROBE_SIGNED` — accumulate the SIGNED activation sum
16866/// instead of its magnitude: what a dropped neuron contributes ON
16867/// AVERAGE, which is the bias a narrowed FFN can add back for free.
16868fn probe_signed() -> bool {
16869    static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16870    *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
16871}
16872
16873/// `CMF_FFN_MEANFILL=<file>` — a masked-out neuron contributes its MEAN
16874/// activation instead of zero (`u32 layers, u32 inter, f32[…]`, the mass
16875/// dump layout, holding per-neuron means). Dropping a neuron outright
16876/// also drops its average contribution, which shifts the layer output by
16877/// a constant; filling the mean back is one add per layer and costs no
16878/// bytes off the bus. This is the measurement arm — in a tube file the
16879/// same correction ships as a per-task bias vector.
16880fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
16881    static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
16882    M.get_or_init(|| {
16883        let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
16884        let b = std::fs::read(&p).ok()?;
16885        let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
16886        let vals: Vec<f32> = b[8..]
16887            .chunks_exact(4)
16888            .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16889            .collect();
16890        eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
16891        Some((inter, vals))
16892    })
16893    .as_ref()
16894}
16895
16896/// `CMF_FFN_PROBE_TOPK` — 0 (default) = accumulate mass, k>0 = count
16897/// how often a neuron lands in a token's top k.
16898fn probe_topk() -> usize {
16899    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16900    *K.get_or_init(|| {
16901        std::env::var("CMF_FFN_PROBE_TOPK")
16902            .ok()
16903            .and_then(|v| v.parse().ok())
16904            .unwrap_or(0)
16905    })
16906}
16907
16908thread_local! {
16909    /// DTG-MA activation probe: per-layer per-neuron Σ|silu(g)·u|
16910    /// accumulator, alive only during `Pipeline::probe_ffn_mass`.
16911    static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
16912        const { std::cell::RefCell::new(None) };
16913}
16914
16915/// Per-token structured sparsity, paid for in bytes.
16916///
16917/// The gate is the cheapest third of an FFN and it already says which
16918/// neurons matter: `silu(gate)` near zero means the neuron contributes
16919/// nothing whatever `up` says. So compute every gate, keep the `k`
16920/// loudest, and read ONLY those neurons' `up` rows and `down` rows —
16921/// the latter needs `down_proj` stored transposed, otherwise a neuron's
16922/// down weights are a strided column and "reading only those" costs a
16923/// full cache line each.
16924///
16925/// Returns `None` when the file has no transposed `down` (the caller
16926/// then runs the ordinary dense path).
16927fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
16928    // The scatter path reads individual rows/columns and cannot express the
16929    // per-matrix signed FWHT boundary.  Let the descriptor-aware dense path
16930    // handle Prism files rather than silently running an unrotated sparse
16931    // approximation.
16932    if d.gate_proj.has_prism_contract()
16933        || d.up_proj.has_prism_contract()
16934        || d.down_proj.has_prism_contract()
16935    {
16936        return None;
16937    }
16938    let dt = d.down_t.as_ref()?;
16939    let inter = d.gate_proj.rows();
16940    let hidden = dt.cols();
16941    if k == 0 || k >= inter || d.act != Act::Silu {
16942        return None;
16943    }
16944    DYN_SCRATCH.with(|sc| {
16945        let mut sc = sc.borrow_mut();
16946        let DynScratch {
16947            g,
16948            mag,
16949            live,
16950            parts,
16951        } = &mut *sc;
16952        g.resize(inter, 0.0);
16953        d.gate_proj.matvec(x, g, pool);
16954        for v in g.iter_mut() {
16955            *v = inference::silu(*v);
16956        }
16957        // The k-th largest |silu(gate)| is the threshold; ties keep more,
16958        // which is the safe side.
16959        mag.clear();
16960        mag.extend(g.iter().map(|v| v.abs()));
16961        let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16962            b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16963        });
16964        let thr = *kth;
16965        live.clear();
16966        live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
16967        let mut out = vec![0.0f32; hidden];
16968        match pool {
16969            Some(p) if live.len() >= 64 => {
16970                let nw = p.n_workers() + 1;
16971                parts.clear();
16972                parts.resize(nw * hidden, 0.0);
16973                let ptr = SendMut(parts.as_mut_ptr());
16974                let n = live.len();
16975                let live_ref: &[u32] = live;
16976                let g_ref: &[f32] = g;
16977                p.run(&|w, workers| {
16978                    let chunk = n.div_ceil(workers);
16979                    let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
16980                    if s >= e {
16981                        return;
16982                    }
16983                    WORKER_SCRATCH.with(|ws| {
16984                        let mut ws = ws.borrow_mut();
16985                        let [scratch, acc] = &mut *ws;
16986                        scratch.resize(hidden.max(x.len()), 0.0);
16987                        acc.clear();
16988                        acc.resize(hidden, 0.0);
16989                        for (o, &nrm) in live_ref[s..e].iter().enumerate() {
16990                            // One neuron of runway: the next row's lines
16991                            // start moving while this one is multiplied.
16992                            if let Some(&nx) = live_ref[s..e].get(o + 1) {
16993                                d.up_proj.prefetch_row(nx as usize);
16994                                dt.prefetch_row(nx as usize);
16995                            }
16996                            let idx = nrm as usize;
16997                            let up = d.up_proj.row_dot(idx, x, scratch);
16998                            let a = g_ref[idx] * up;
16999                            if a != 0.0 {
17000                                dt.add_row_scaled(idx, a, acc, scratch);
17001                            }
17002                        }
17003                        for (j, v) in acc.iter().enumerate() {
17004                            unsafe { *ptr.at(w * hidden + j) = *v };
17005                        }
17006                    });
17007                });
17008                for w in 0..nw {
17009                    for (j, o) in out.iter_mut().enumerate() {
17010                        *o += parts[w * hidden + j];
17011                    }
17012                }
17013            }
17014            _ => {
17015                WORKER_SCRATCH.with(|ws| {
17016                    let mut ws = ws.borrow_mut();
17017                    let [scratch, _acc] = &mut *ws;
17018                    scratch.resize(hidden.max(x.len()), 0.0);
17019                    for &nrm in live.iter() {
17020                        let idx = nrm as usize;
17021                        let up = d.up_proj.row_dot(idx, x, scratch);
17022                        let a = g[idx] * up;
17023                        if a != 0.0 {
17024                            dt.add_row_scaled(idx, a, &mut out, scratch);
17025                        }
17026                    }
17027                });
17028            }
17029        }
17030        Some(out)
17031    })
17032}
17033
17034/// Caller-side scratch of the dynamic path — one allocation per thread,
17035/// not one per layer per token (that alone cost a third of the decode).
17036struct DynScratch {
17037    g: Vec<f32>,
17038    mag: Vec<f32>,
17039    live: Vec<u32>,
17040    parts: Vec<f32>,
17041}
17042
17043thread_local! {
17044    static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
17045        std::cell::RefCell::new(DynScratch {
17046            g: Vec::new(),
17047            mag: Vec::new(),
17048            live: Vec::new(),
17049            parts: Vec::new(),
17050        })
17051    };
17052    /// Pool-worker scratch: the row buffer and this worker's partial sum.
17053    static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
17054        const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
17055}
17056
17057/// `dense_ffn_cpu` with a per-visit mask landing on the activations —
17058/// the masked-inference fast path's decode arm. Full fused quant
17059/// compute, closed neurons zeroed before down: arithmetically the
17060/// pruned network, no dequant, no weight bytes touched.
17061fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
17062    let inter = d.gate_proj.rows();
17063    FFN_SCRATCH.with(|s| {
17064        let mut s = s.borrow_mut();
17065        let [g, u, ..] = &mut *s;
17066        g.resize(inter, 0.0);
17067        if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
17068            // g holds silu(gate)·up.
17069        } else {
17070            u.resize(inter, 0.0);
17071            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
17072            for i in 0..inter {
17073                g[i] = d.act.combine(g[i], u[i]);
17074            }
17075        }
17076        zero_masked_cols(g, 1, inter, mask_row);
17077        let mut out = attention::take_buf(d.down_proj.rows());
17078        d.down_proj.matvec(g, &mut out, pool);
17079        out
17080    })
17081}
17082
17083/// Dense FFN as one GPU submission via the MoE block path (single
17084/// expert, weight 1.0): gate → silu·up → down chained in one command
17085/// buffer, intermediate activations device-resident. None → weights
17086/// not q8-mapped in the primary shard / over the VRAM budget / backend
17087/// refusal → honest CPU path.
17088fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
17089    if d.gate_proj.has_prism_contract()
17090        || d.up_proj.has_prism_contract()
17091        || d.down_proj.has_prism_contract()
17092    {
17093        return None;
17094    }
17095    // The GPU block hardcodes SiLU; GeLU FFNs (Gemma) stay on CPU.
17096    if d.act != Act::Silu {
17097        return None;
17098    }
17099    // Threshold: tiny FFNs are not worth a submission (q1 excepted —
17100    // see the caller's gate).
17101    if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
17102        return None;
17103    }
17104    let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
17105    let mut model_ref = None;
17106    moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
17107    let model = model_ref?;
17108    let hidden = jobs[0].down.1;
17109    let mut out = attention::take_buf(hidden);
17110    if crate::gpu::moe_block(&model, &jobs, &mut out) {
17111        Some(out)
17112    } else {
17113        let mut out = out;
17114        attention::recycle_buf(&mut out);
17115        None
17116    }
17117}
17118
17119/// q8-mapped primary-shard tensor parts for a GPU job: q8_2f carries
17120/// its column field, q8_row runs with empty col slices (the backend
17121/// skips the multiply). Shared by the MoE block and the dense-FFN
17122/// single-job path.
17123#[allow(clippy::type_complexity)]
17124#[allow(clippy::type_complexity)]
17125pub(crate) fn moe_parts(
17126    t: &QTensor,
17127) -> Option<(
17128    &std::sync::Arc<cortiq_core::CmfModel>,
17129    usize,
17130    usize,
17131    usize,
17132    &[f32],
17133    &[f32],
17134    bool,
17135    bool,
17136    bool,
17137)> {
17138    match t {
17139        QTensor::Mapped {
17140            model,
17141            idx,
17142            dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
17143            rows,
17144            cols,
17145            row_scale,
17146            col_field,
17147            ..
17148        } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
17149            model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
17150        )),
17151        // q1: tile-embedded scales — empty rs/col slices, raw xs.
17152        QTensor::Mapped {
17153            model,
17154            idx,
17155            dtype: cortiq_core::TensorDtype::Q1,
17156            rows,
17157            cols,
17158            ..
17159        } => Some((
17160            model,
17161            *idx,
17162            *rows,
17163            *cols,
17164            &[][..],
17165            &[][..],
17166            true,
17167            false,
17168            false,
17169        )),
17170        // q4_tiled: 18-byte tiles with embedded f16 scales — raw xs.
17171        QTensor::Mapped {
17172            model,
17173            idx,
17174            dtype: cortiq_core::TensorDtype::Q4Tiled,
17175            rows,
17176            cols,
17177            ..
17178        } => Some((
17179            model,
17180            *idx,
17181            *rows,
17182            *cols,
17183            &[][..],
17184            &[][..],
17185            false,
17186            true,
17187            false,
17188        )),
17189        // q4tp: same raw-xs contract, different stride and scale plane.
17190        QTensor::Mapped {
17191            model,
17192            idx,
17193            dtype: cortiq_core::TensorDtype::Q4TiledP,
17194            rows,
17195            cols,
17196            ..
17197        } => Some((
17198            model,
17199            *idx,
17200            *rows,
17201            *cols,
17202            &[][..],
17203            &[][..],
17204            false,
17205            true,
17206            false,
17207        )),
17208        // q2tp: the 2-bit expert plane of the mixed profile — q4 family
17209        // for stride bookkeeping, flagged q2 so the trio validation can
17210        // demand a q4tp down.
17211        QTensor::Mapped {
17212            model,
17213            idx,
17214            dtype: cortiq_core::TensorDtype::Q2TiledP,
17215            rows,
17216            cols,
17217            ..
17218        } => Some((
17219            model,
17220            *idx,
17221            *rows,
17222            *cols,
17223            &[][..],
17224            &[][..],
17225            false,
17226            true,
17227            true,
17228        )),
17229        _ => None,
17230    }
17231}
17232
17233/// Map a MoE onto the Metal token graph's contract: f32 router, a
17234/// shared expert (gated — Qwen — or ungated at weight 1 — DeepSeek-V3 /
17235/// HunYuan hy_v3), softmax or sigmoid scores with an optional selection
17236/// bias and routed scale, experts uniformly q4tp (or the mixed profile:
17237/// q2tp gate/up over a q4tp down). τ routers, masks, per-expert scales
17238/// and Gemma's router-input norm refuse here — those semantics stay on
17239/// the CPU path.
17240#[cfg(target_os = "macos")]
17241fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
17242    if m.router_input_norm
17243        || m.route_tau.is_some()
17244        || m.mask.is_some()
17245        || m.per_expert_scale.is_some()
17246        || m.experts.is_empty()
17247        || m.top_k == 0
17248        || m.resonance.is_some()
17249    {
17250        return None;
17251    }
17252    // The select kernel always fills the shared slot: a model without a
17253    // shared expert (LFM2-MoE) stays on the CPU path here.
17254    let (sh, sg) = match &m.shared {
17255        Some((sh, sg)) => (sh, sg.as_ref()),
17256        None => return None,
17257    };
17258    let (rf, rr, rc) = m.router.f32_parts()?;
17259    if rr != m.experts.len() || rc != hidden {
17260        return None;
17261    }
17262    let shared_gated = sg.is_some();
17263    let sf = match sg {
17264        Some(sg) => {
17265            let (sf, sr, sc) = sg.f32_parts()?;
17266            if sr * sc != hidden {
17267                return None;
17268            }
17269            sf
17270        }
17271        // Ungated: the router's first row stands in for the gate matvec
17272        // (its logit is never read — the kernel pins weight 1).
17273        None => &rf[..hidden],
17274    };
17275    if let Some(b) = &m.expert_bias {
17276        if b.len() != m.experts.len() {
17277            return None;
17278        }
17279    }
17280    let inter = m.experts[0].gate_proj.rows();
17281    // The first expert's gate decides the profile; every trio (shared
17282    // included) must agree — the jobs ladder flips ONE kernel for all.
17283    let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
17284    let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
17285        if e.act != Act::Silu
17286            || e.gate_proj.rows() != inter
17287            || e.gate_proj.cols() != hidden
17288            || e.up_proj.rows() != inter
17289            || e.up_proj.cols() != hidden
17290            || e.down_proj.rows() != hidden
17291            || e.down_proj.cols() != inter
17292        {
17293            return None;
17294        }
17295        let pick = |t: &QTensor| -> Option<usize> {
17296            if gu_q2 {
17297                t.mapped_q2tp().map(|(_, i)| i)
17298            } else {
17299                t.mapped_q4tp().map(|(_, i)| i)
17300            }
17301        };
17302        Some((
17303            pick(&e.gate_proj)?,
17304            pick(&e.up_proj)?,
17305            e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
17306        ))
17307    };
17308    let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
17309    let shared = trio(sh)?;
17310    Some(crate::gpu::GpuMoe {
17311        router: rf,
17312        sgate: sf,
17313        experts,
17314        shared,
17315        n_exp: m.experts.len(),
17316        top_k: m.top_k,
17317        inter,
17318        norm_topk: m.norm_topk_prob,
17319        route_scale: m.routed_scaling,
17320        gu_q2,
17321        sigmoid: m.router_sigmoid,
17322        bias: m.expert_bias.as_deref(),
17323        shared_gated,
17324    })
17325}
17326
17327/// Build one gate/up/down GPU job from three tensors. `moe_push_job` is the
17328/// DenseFfn-shaped caller; architectures that keep their experts in their own
17329/// structs (DeepSeek-V4) come here directly.
17330pub(crate) fn moe_push_job_parts<'a>(
17331    gate: &'a QTensor,
17332    up: &'a QTensor,
17333    down: &'a QTensor,
17334    x: &[f32],
17335    w: f32,
17336    swiglu_limit: f32,
17337    jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17338    model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17339) -> Option<()> {
17340    use crate::qtensor::prescale;
17341    let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
17342    let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
17343    let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
17344    if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17345        return None; // mixed-dtype trio — honest CPU path
17346    }
17347    // The 2-bit profile is gate/up q2tp over a PLAIN q4tp down; any other
17348    // 2-bit arrangement stays on the CPU.
17349    if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
17350        return None;
17351    }
17352    if !gq2 && dq2 {
17353        return None;
17354    }
17355    model_ref.get_or_insert_with(|| gm.clone());
17356    let dt = |cf: &[f32]| {
17357        if cf.is_empty() {
17358            cortiq_core::TensorDtype::Q8Row
17359        } else {
17360            cortiq_core::TensorDtype::Q8_2f
17361        }
17362    };
17363    jobs.push(crate::gpu::MoeJob {
17364        gate: (gi, gr, gc, grs),
17365        up: (ui, ur, uc, urs),
17366        down: (di, dr, dc, drs),
17367        xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
17368        xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
17369        down_col: dcf,
17370        w,
17371        q1: gq1,
17372        q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
17373        q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
17374        gu_q2: gq2,
17375        swiglu_limit,
17376    });
17377    Some(())
17378}
17379
17380/// Build one gate/up/down GPU job (see `moe_parts`).
17381fn moe_push_job<'a>(
17382    d: &'a DenseFfn,
17383    x: &[f32],
17384    w: f32,
17385    jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17386    model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17387) -> Option<()> {
17388    use crate::qtensor::prescale;
17389    if d.act != Act::Silu {
17390        return None; // GPU block hardcodes SiLU
17391    }
17392    let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
17393    let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
17394    let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
17395    if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17396        return None; // mixed-dtype trio — honest CPU path
17397    }
17398    if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
17399        return None;
17400    }
17401    if !gq2 && dq2 {
17402        return None;
17403    }
17404    model_ref.get_or_insert_with(|| gm.clone());
17405    let gdt = if gcf.is_empty() {
17406        cortiq_core::TensorDtype::Q8Row
17407    } else {
17408        cortiq_core::TensorDtype::Q8_2f
17409    };
17410    let udt = if ucf.is_empty() {
17411        cortiq_core::TensorDtype::Q8Row
17412    } else {
17413        cortiq_core::TensorDtype::Q8_2f
17414    };
17415    jobs.push(crate::gpu::MoeJob {
17416        gate: (gi, gr, gc, grs),
17417        up: (ui, ur, uc, urs),
17418        down: (di, dr, dc, drs),
17419        xs_gate: prescale(x, gcf, gdt).into_owned(),
17420        xs_up: prescale(x, ucf, udt).into_owned(),
17421        down_col: dcf,
17422        w,
17423        q1: gq1,
17424        q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
17425        q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
17426        gu_q2: gq2,
17427        swiglu_limit: 0.0,
17428    });
17429    Some(())
17430}
17431
17432/// Sparse dense-FFN directly on QUANTIZED weights (mask × mmap): reads
17433/// ONLY the active neurons' gate/up rows and down columns from the mmap
17434/// — no full-matrix dequant, no f32 model copy. This is what lets a
17435/// masked big model run at quantized RSS (the historical mask path
17436/// forced the whole model to f32). Semantics identical to the f32
17437/// sparse path within quant tolerance.
17438fn sparse_ffn_quant(
17439    d: &DenseFfn,
17440    x: &[f32],
17441    active: &[u16],
17442    hidden: usize,
17443    pool: Option<&Pool>,
17444) -> Vec<f32> {
17445    let n = active.len();
17446    let inter = d.gate_proj.rows();
17447    let mut act = vec![0.0f32; n];
17448    // Scratch is needed if EITHER projection is group-packed (q4/vbit);
17449    // gate/up normally share a dtype but sizing on both is robust.
17450    let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
17451    let compute = |ai: usize| -> f32 {
17452        let idx = active[ai] as usize;
17453        if idx >= inter {
17454            return 0.0; // defensive parity with the f32 sparse path
17455        }
17456        let mut s = if need_scratch {
17457            vec![0.0f32; hidden]
17458        } else {
17459            Vec::new()
17460        };
17461        let gate = d.gate_proj.row_dot(idx, x, &mut s);
17462        let up = d.up_proj.row_dot(idx, x, &mut s);
17463        d.act.combine(gate, up)
17464    };
17465    match pool {
17466        Some(p) if n >= 256 => {
17467            let ptr = SendMut(act.as_mut_ptr());
17468            p.run(&|widx, nw| {
17469                let chunk = n.div_ceil(nw);
17470                let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
17471                for ai in s..e {
17472                    unsafe { *ptr.at(ai) = compute(ai) };
17473                }
17474            });
17475        }
17476        _ => {
17477            for (ai, a) in act.iter_mut().enumerate() {
17478                *a = compute(ai);
17479            }
17480        }
17481    }
17482    // Scatter through active down columns (reads only those columns).
17483    let mut out = vec![0.0f32; hidden];
17484    for (ai, &idx) in active.iter().enumerate() {
17485        let w = act[ai];
17486        if w.abs() >= 1e-12 && (idx as usize) < inter {
17487            d.down_proj.add_col_scaled(idx as usize, w, &mut out);
17488        }
17489    }
17490    out
17491}
17492
17493/// Test-only re-export of the private sparse-quant FFN (mask × mmap gate).
17494#[doc(hidden)]
17495pub fn sparse_ffn_quant_for_test(
17496    d: &DenseFfn,
17497    x: &[f32],
17498    active: &[u16],
17499    hidden: usize,
17500) -> Vec<f32> {
17501    sparse_ffn_quant(d, x, active, hidden, None)
17502}
17503
17504/// Dequantize a DenseFfn's three matrices to f32 (transient; only the
17505/// q4/vbit-masked fallback uses it — the memory-lean path is
17506/// sparse_ffn_quant). Reuses row_f32 row-by-row.
17507fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
17508    let deq = |t: &QTensor| -> Vec<f32> {
17509        let (rows, cols) = (t.rows(), t.cols());
17510        let mut out = vec![0.0f32; rows * cols];
17511        for r in 0..rows {
17512            t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
17513        }
17514        out
17515    };
17516    (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
17517}
17518
17519/// Pointer wrapper for the worker-pool scatter (same pattern as qtensor).
17520struct SendMut(*mut f32);
17521unsafe impl Send for SendMut {}
17522unsafe impl Sync for SendMut {}
17523impl SendMut {
17524    #[inline]
17525    // Deliberate unsynchronized scatter: pool workers write disjoint indices
17526    // in parallel, so returning `&mut` from `&self` is intentional here.
17527    #[allow(clippy::mut_from_ref)]
17528    unsafe fn at(&self, i: usize) -> &mut f32 {
17529        unsafe { &mut *self.0.add(i) }
17530    }
17531}
17532
17533/// Router → (selected experts in torch.topk order, per-expert score
17534/// vector, normalizer). The final weight of expert `e` is `p[e] / wsum`.
17535///
17536/// Two regimes share this. Qwen: softmax over ALL experts, top-k of the
17537/// probabilities, optional renorm — `router_sigmoid=false`, no bias,
17538/// scale 1 → bit-identical to the historical path. LFM2-MoE /
17539/// DeepSeek-V3 `noaux_tc`: per-expert sigmoid scores, an optional
17540/// selection bias (top-k CHOICE only; weights stay unbiased), a 1e-6 renorm
17541/// floor and a routed scale. Architectures whose reference uses a different
17542/// sigmoid denominator floor (for example GLM-5's `1e-20`) call
17543/// [`moe_route_with_eps`] directly; the historical generic path remains
17544/// unchanged.
17545pub(crate) fn moe_route(
17546    logits: &[f32],
17547    m: &MoeFfn,
17548    allowed: Option<&[bool]>,
17549) -> (Vec<usize>, Vec<f32>, f32) {
17550    moe_route_with_eps(logits, m, allowed, 1e-6)
17551}
17552
17553/// Router implementation with an explicit sigmoid renormalization floor.
17554///
17555/// GLM-5.3's source computes `sum(selected_scores) + 1e-20`; using the
17556/// generic 1e-6 floor there is not a harmless tolerance difference when all
17557/// logits are very negative: it collapses the routed branch toward zero
17558/// instead of normalizing the selected experts. Keeping the epsilon parameter
17559/// here avoids changing the established Qwen/LFM2 contract while allowing
17560/// each architecture to preserve its own numerical semantics.
17561pub(crate) fn moe_route_with_eps(
17562    logits: &[f32],
17563    m: &MoeFfn,
17564    allowed: Option<&[bool]>,
17565    sigmoid_denom_eps: f32,
17566) -> (Vec<usize>, Vec<f32>, f32) {
17567    let ne = logits.len();
17568    // Expert restriction: the static env mask (CMF_MOE_MASK) AND the
17569    // active task mask's expert fields (spec §5) both narrow the
17570    // candidate set; selection happens over the admitted experts only.
17571    // With norm_topk the kept weights renormalize below; without it
17572    // the excluded mass is honestly dropped.
17573    let admit = |e: usize| {
17574        m.mask.as_ref().is_none_or(|mk| mk[e])
17575            && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
17576    };
17577    // The resonance router (spec §9.5.1) selects by the RAW score: the
17578    // trainer (`resonance_winner`) and the resident graph
17579    // (`embryo_core_route_pick`) take the first maximum of the scores
17580    // and run the winner with weight 1.0. Selecting through the softmax
17581    // instead is not the same decision: `exp(l − max)` rounds two scores
17582    // closer than 2^-25 (possible below |score| 0.25) to the same 1.0, and
17583    // the lower index would take a token whose score is strictly smaller
17584    // — the trainer's trace and the graph would disagree with this path.
17585    // `−∞` (outside the shell) never wins; with no finite admitted expert
17586    // the generic path below degrades to uniform. The winner's
17587    // probability is 1.0 by construction (a one-hot `p`), so its
17588    // renormalized weight is `routed_scaling` on both norm_topk settings.
17589    if m.resonance.is_some() && m.top_k == 1 {
17590        let mut best: Option<usize> = None;
17591        for e in (0..ne).filter(|&e| admit(e)) {
17592            let l = logits[e];
17593            if l == f32::NEG_INFINITY || l.is_nan() {
17594                continue;
17595            }
17596            if best.is_none_or(|b| l > logits[b]) {
17597                best = Some(e);
17598            }
17599        }
17600        if let Some(b) = best {
17601            let mut p = vec![0.0f32; ne];
17602            p[b] = 1.0;
17603            return (vec![b], p, 1.0 / m.routed_scaling);
17604        }
17605    }
17606    // A `−∞` logit (a grown expert outside its shell, `Resonance::scores`)
17607    // takes probability 0 on both paths: sigmoid(−∞) = 0, exp(−∞ − max) = 0
17608    // — top-1 is the best FINITE expert, its renormalized weight exactly
17609    // 1.0. Every expert at −∞ cannot happen (trunk experts have no shell);
17610    // should it, the softmax would be NaN, so it degrades to uniform.
17611    let p: Vec<f32> = if m.router_sigmoid {
17612        logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
17613    } else {
17614        let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
17615        if mx == f32::NEG_INFINITY {
17616            vec![1.0 / ne.max(1) as f32; ne]
17617        } else {
17618            let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
17619            let s: f32 = e.iter().sum();
17620            for v in &mut e {
17621                *v /= s;
17622            }
17623            e
17624        }
17625    };
17626    let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
17627    // Descending by selection score, lower index wins ties (torch.topk).
17628    match &m.expert_bias {
17629        Some(b) => idx.sort_unstable_by(|&x, &y| {
17630            (p[y] + b[y])
17631                .partial_cmp(&(p[x] + b[x]))
17632                .unwrap()
17633                .then(x.cmp(&y))
17634        }),
17635        None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
17636    }
17637    idx.truncate(m.top_k);
17638    // Adaptive τ-routing: trim the tail experts once the kept mass is
17639    // enough. wsum below renormalizes over the KEPT set, so the output
17640    // stays a proper weighted average.
17641    if let Some(tau) = m.route_tau {
17642        let total: f32 = idx.iter().map(|&e| p[e]).sum();
17643        if total > 0.0 {
17644            let mut acc = 0.0f32;
17645            let mut keep = idx.len();
17646            for (i, &e) in idx.iter().enumerate() {
17647                acc += p[e];
17648                if acc >= tau * total {
17649                    keep = i + 1;
17650                    break;
17651                }
17652            }
17653            idx.truncate(keep);
17654        }
17655    }
17656    let wsum: f32 = if m.norm_topk_prob {
17657        let s: f32 = idx.iter().map(|&e| p[e]).sum();
17658        // Sigmoid routers use their architecture's reference floor; the
17659        // softmax path's probs already sum near 1, so it stays exactly as
17660        // before.
17661        (if m.router_sigmoid {
17662            s + sigmoid_denom_eps
17663        } else {
17664            s
17665        }) / m.routed_scaling
17666    } else {
17667        1.0 / m.routed_scaling
17668    };
17669    (idx, p, wsum)
17670}
17671
17672/// See the call site: one `layer:e1,e2,…` line per routed token.
17673fn moe_trace(idx: &[usize]) {
17674    moe_trace_at(crate::gpu::cur_layer() as i32, idx)
17675}
17676
17677/// The same, for callers that know their layer (DSV4 owns its layers and
17678/// never sets the pipeline's current-layer marker).
17679pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
17680    use std::io::Write;
17681    static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
17682        std::sync::OnceLock::new();
17683    let Some(f) = F.get_or_init(|| {
17684        let p = std::env::var("CMF_MOE_TRACE").ok()?;
17685        Some(std::sync::Mutex::new(
17686            std::fs::OpenOptions::new()
17687                .create(true)
17688                .append(true)
17689                .open(p)
17690                .ok()?,
17691        ))
17692    }) else {
17693        return;
17694    };
17695    let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
17696    let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
17697}
17698
17699/// MoE FFN: router → top-k experts (see `moe_route`). Only selected
17700/// experts' pages are touched in mmap.
17701pub(crate) fn moe_ffn(
17702    m: &MoeFfn,
17703    x: &[f32],
17704    pool: Option<&Pool>,
17705    allowed: Option<&[bool]>,
17706) -> Vec<f32> {
17707    let r = moe_ffn_route(m, x, pool, allowed);
17708    moe_ffn_experts(m, x, &r, pool)
17709}
17710
17711/// One token's host route through a MoE layer: the chosen experts in
17712/// selection order, the per-expert scores and the normalizer (see
17713/// `moe_route`), plus the raw router logits.
17714pub(crate) struct MoeRoute {
17715    pub idx: Vec<usize>,
17716    pub p: Vec<f32>,
17717    pub wsum: f32,
17718    pub logits: Vec<f32>,
17719}
17720
17721/// The routing half of `moe_ffn`, shared by every executor of the chosen
17722/// experts (the host/per-op path below and the MiMo dynamic device cache,
17723/// `crate::mimo_moe`): activation accounting, router logits, `moe_route`,
17724/// the selection statistics and the `CMF_MOE_TRACE` line — so switching
17725/// executors can never change which experts a token gets.
17726pub(crate) fn moe_ffn_route(
17727    m: &MoeFfn,
17728    x: &[f32],
17729    pool: Option<&Pool>,
17730    allowed: Option<&[bool]>,
17731) -> MoeRoute {
17732    accumulate_act(m, x, 1);
17733    let ne = m.experts.len();
17734    let mut logits = vec![0.0f32; ne];
17735    match &m.resonance {
17736        Some(r) => r.scores(x, &mut logits),
17737        None => m.router.matvec(x, &mut logits, pool),
17738    }
17739    let (idx, p, wsum) = moe_route(&logits, m, allowed);
17740    {
17741        let mut st = m.stats.borrow_mut();
17742        if st.len() < ne {
17743            st.resize(ne, 0);
17744        }
17745        for &e in &idx {
17746            st[e] += 1;
17747        }
17748    }
17749    // `CMF_MOE_TRACE=<file>`: append one line per (layer, token) with the
17750    // selected expert ids. The cumulative `stats` above answer "which
17751    // experts are popular"; a residency design needs the question they
17752    // cannot answer — whether CONSECUTIVE tokens reuse experts (the
17753    // temporal locality an LRU cache lives on, FreeToken §4).
17754    moe_trace(&idx);
17755    MoeRoute {
17756        idx,
17757        p,
17758        wsum,
17759        logits,
17760    }
17761}
17762
17763/// The expert half of `moe_ffn`: run a route's experts on the per-op GPU
17764/// block or the host.
17765pub(crate) fn moe_ffn_experts(
17766    m: &MoeFfn,
17767    x: &[f32],
17768    r: &MoeRoute,
17769    pool: Option<&Pool>,
17770) -> Vec<f32> {
17771    let (idx, p, wsum) = (&r.idx, &r.p, r.wsum);
17772    // D5: the whole layer MoE block in one GPU command buffer (experts — the
17773    // same mmap via a no-copy buffer; intermediate activations on the GPU).
17774    // Same Ffn probe class as the dense chain: one submit per layer
17775    // either wins on this driver stack or it doesn't.
17776    if crate::gpu::enabled_here() {
17777        match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
17778            crate::gpu::ProbeArm::Gpu => {
17779                let t0 = std::time::Instant::now();
17780                if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
17781                    crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
17782                    return out;
17783                }
17784            }
17785            crate::gpu::ProbeArm::CpuTimed => {
17786                let t0 = std::time::Instant::now();
17787                let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
17788                crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
17789                return out;
17790            }
17791            crate::gpu::ProbeArm::Cpu => {
17792                return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
17793            }
17794        }
17795    }
17796    moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
17797}
17798
17799/// One MoE token through the MiMo expert bank (`crate::mimo_moe`), or —
17800/// when the bank does not serve it — through the host path with the SAME
17801/// route, so the routing statistics and `CMF_MOE_TRACE` see it once.
17802fn moe_ffn_banked(
17803    slot: &mut crate::mimo_moe::Slot,
17804    li: usize,
17805    m: &MoeFfn,
17806    x: &[f32],
17807    pool: Option<&Pool>,
17808) -> Vec<f32> {
17809    let t0 = std::time::Instant::now();
17810    let r = moe_ffn_route(m, x, pool, None);
17811    slot.note_route(t0.elapsed().as_nanos() as u64);
17812    match slot.forward(li, m, x, &r, pool) {
17813        Some(out) => out,
17814        None => crate::qtensor::float_activations_scope(|| {
17815            crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, &r, pool))
17816        }),
17817    }
17818}
17819
17820/// Verify rows share a bank frame; routing and fallback are decode's.
17821fn moe_ffn_banked_rows(
17822    slot: &mut crate::mimo_moe::Slot,
17823    li: usize,
17824    m: &MoeFfn,
17825    xs: &[f32],
17826    b: usize,
17827    hidden: usize,
17828    pool: Option<&Pool>,
17829) -> Vec<f32> {
17830    let t0 = std::time::Instant::now();
17831    let routes: Vec<_> = xs
17832        .chunks_exact(hidden)
17833        .map(|x| moe_ffn_route(m, x, pool, None))
17834        .collect();
17835    slot.note_route(t0.elapsed().as_nanos() as u64);
17836    if let Some(out) = slot.forward_rows(li, m, xs, &routes, pool) {
17837        return out;
17838    }
17839    let mut out = Vec::with_capacity(b * hidden);
17840    for (x, r) in xs.chunks_exact(hidden).zip(&routes) {
17841        let row = slot.forward(li, m, x, r, pool).unwrap_or_else(|| {
17842            // A failed bank must not stream missing experts into the arena.
17843            crate::qtensor::float_activations_scope(|| {
17844                crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, r, pool))
17845            })
17846        });
17847        out.extend(row);
17848    }
17849    out
17850}
17851
17852/// One-shot report of whether the whole-token wgpu graph actually formed.
17853/// A refusal silently reverts to the per-op path, which is how a model can
17854/// look "GPU-accelerated" while every layer walks the host.  A device prefix
17855/// is tracked separately because it still pays a host boundary for the tail.
17856fn graph_note(built: bool, layers_run: usize, total_layers: usize) {
17857    use std::sync::atomic::{AtomicBool, Ordering};
17858    if built {
17859        GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
17860        if total_layers > 0 && layers_run < total_layers {
17861            GRAPH_TOK_PREFIX.fetch_add(1, Ordering::Relaxed);
17862        } else {
17863            GRAPH_TOK_FULL.fetch_add(1, Ordering::Relaxed);
17864        }
17865    } else {
17866        GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
17867    }
17868    static SAID: AtomicBool = AtomicBool::new(false);
17869    if !SAID.swap(true, Ordering::Relaxed) {
17870        if built {
17871            tracing::info!("wgpu whole-token graph: ACTIVE");
17872        } else {
17873            tracing::warn!("wgpu whole-token graph refused — per-op path");
17874        }
17875    }
17876}
17877
17878/// Whole-token graph outcomes, process-wide: a benchmark that claims a
17879/// GPU number while MISS climbs is measuring the CPU — the honest-bench
17880/// contract makes that an error, not a footnote.
17881pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17882pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17883/// Graph calls that returned a hidden after running only a leading device
17884/// prefix.  These are valid hybrid executions but must not be reported as a
17885/// full GPU graph in benchmark evidence.
17886pub static GRAPH_TOK_PREFIX: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17887/// Graph calls that covered the complete requested layer span.
17888pub static GRAPH_TOK_FULL: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17889
17890/// Native Metal TokenGraph completion counters. These are incremented only
17891/// after checked command-buffer completion and successful readback, so a
17892/// fused-head NLL report can prove the route rather than infer it from env.
17893pub static METAL_GRAPH_TOK_OK: std::sync::atomic::AtomicU64 =
17894    std::sync::atomic::AtomicU64::new(0);
17895pub static METAL_GRAPH_HEAD_OK: std::sync::atomic::AtomicU64 =
17896    std::sync::atomic::AtomicU64::new(0);
17897pub static METAL_GRAPH_HEAD_MISS: std::sync::atomic::AtomicU64 =
17898    std::sync::atomic::AtomicU64::new(0);
17899pub static METAL_GRAPH_LAYERS: std::sync::atomic::AtomicU64 =
17900    std::sync::atomic::AtomicU64::new(0);
17901pub static METAL_GRAPH_ERRORS: std::sync::atomic::AtomicU64 =
17902    std::sync::atomic::AtomicU64::new(0);
17903/// Ordinary native-Metal rows-prefill admissions and completed rows.  These
17904/// counters are separate from TokenGraph token/head counts so a batch NLL
17905/// receipt cannot accidentally claim serial execution as batched.
17906pub static METAL_PREFILL_CHUNKS: std::sync::atomic::AtomicU64 =
17907    std::sync::atomic::AtomicU64::new(0);
17908pub static METAL_PREFILL_ROWS: std::sync::atomic::AtomicU64 =
17909    std::sync::atomic::AtomicU64::new(0);
17910pub static METAL_PREFILL_HEAD_ROWS: std::sync::atomic::AtomicU64 =
17911    std::sync::atomic::AtomicU64::new(0);
17912pub static METAL_PREFILL_ERRORS: std::sync::atomic::AtomicU64 =
17913    std::sync::atomic::AtomicU64::new(0);
17914
17915/// `CMF_MOE_BATCH=0` restores the per-expert serial loop — the A/B lever
17916/// for the batched kernel, and how its bit-identity is checked.
17917fn moe_batch_enabled() -> bool {
17918    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17919    *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
17920}
17921
17922/// Two-dispatch CPU MoE: every routed expert (and the shared one) fused
17923/// into one gate/up/SiLU dispatch and one down dispatch, instead of two
17924/// pool barriers per expert. Bit-identical to the serial loop below —
17925/// see `moe_gate_up_many` / `moe_down_many`. `None` = the batched kernel
17926/// does not cover this layer, walk the serial path.
17927fn moe_ffn_cpu_batched(
17928    m: &MoeFfn,
17929    x: &[f32],
17930    idx: &[usize],
17931    p: &[f32],
17932    wsum: f32,
17933    pool: Option<&Pool>,
17934) -> Option<Vec<f32>> {
17935    if idx.is_empty() || !moe_batch_enabled() {
17936        return None;
17937    }
17938    // The bake probe reads per-neuron activation mass out of the
17939    // single-expert path; batching would skip it. Rare and offline —
17940    // hand those runs to the serial loop.
17941    if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
17942        return None;
17943    }
17944    let n = idx.len() + usize::from(m.shared.is_some());
17945    let mut pairs = Vec::with_capacity(n);
17946    let mut downs = Vec::with_capacity(n);
17947    let mut ws = Vec::with_capacity(n);
17948    for &e in idx {
17949        let d = &m.experts[e];
17950        if d.act != Act::Silu {
17951            return None;
17952        }
17953        pairs.push((&d.gate_proj, &d.up_proj));
17954        downs.push(&d.down_proj);
17955        ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
17956    }
17957    // The shared expert goes last, matching the serial loop's order —
17958    // the f32 accumulation order is part of the bit-identity claim.
17959    if let Some((se, gate)) = &m.shared {
17960        if se.act != Act::Silu {
17961            return None;
17962        }
17963        let g = gate.as_ref().map_or(1.0, |gate| {
17964            let mut gl = [0.0f32; 1];
17965            gate.matvec(x, &mut gl, pool);
17966            1.0 / (1.0 + (-gl[0]).exp())
17967        });
17968        pairs.push((&se.gate_proj, &se.up_proj));
17969        downs.push(&se.down_proj);
17970        ws.push(g);
17971    }
17972    let inter = pairs[0].0.rows();
17973    let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
17974    if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
17975        return None;
17976    }
17977    let mut out = attention::take_buf(x.len());
17978    if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
17979        attention::recycle_buf(&mut out);
17980        return None;
17981    }
17982    Some(out)
17983}
17984
17985/// Exact CPU completion for the routed experts a dynamic device cache did
17986/// not contain. The weights are already the router's final normalized mix.
17987/// Keeping this independent of `MoeFfn` makes the job `Sync`: its routing
17988/// statistics live in a `RefCell`, while the immutable expert tensors can be
17989/// evaluated safely in parallel with the GPU's resident subset.
17990pub(crate) fn moe_cold_experts_cpu(
17991    experts: &[(&DenseFfn, f32)],
17992    x: &[f32],
17993    pool: Option<&Pool>,
17994) -> Vec<f32> {
17995    let mut out = attention::take_buf(x.len());
17996    if experts.is_empty() {
17997        return out;
17998    }
17999    let pairs: Vec<_> = experts
18000        .iter()
18001        .map(|(e, _)| (&e.gate_proj, &e.up_proj))
18002        .collect();
18003    let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
18004    let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
18005    let inter = experts[0].0.gate_proj.rows();
18006    let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
18007    if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
18008        && QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
18009    {
18010        return out;
18011    }
18012    out.fill(0.0);
18013    for &(expert, weight) in experts {
18014        let mut one = dense_ffn(expert, x, pool);
18015        for (o, v) in out.iter_mut().zip(&one) {
18016            *o += weight * v;
18017        }
18018        attention::recycle_buf(&mut one);
18019    }
18020    out
18021}
18022
18023/// Cold part of a short bank batch. Share each expert's weight stream
18024/// across its tokens, but reduce contributions in each token's route order.
18025/// On an unsupported CPU/layout, retain the single-token cold kernels.
18026pub(crate) fn moe_cold_experts_rows_cpu(
18027    jobs: &[Vec<(&DenseFfn, f32)>],
18028    xs: &[f32],
18029    hidden: usize,
18030    pool: Option<&Pool>,
18031) -> Vec<f32> {
18032    let mut out = vec![0.0; xs.len()];
18033    let mut experts: Vec<&DenseFfn> = Vec::new();
18034    let mut groups: Vec<Vec<usize>> = Vec::new();
18035    let mut terms = vec![Vec::new(); jobs.len()];
18036    for (r, row) in jobs.iter().enumerate() {
18037        for &(e, w) in row {
18038            let g = match experts.iter().position(|&d| std::ptr::eq(d, e)) {
18039                Some(g) => g,
18040                None => {
18041                    experts.push(e);
18042                    groups.push(Vec::new());
18043                    groups.len() - 1
18044                }
18045            };
18046            terms[r].push((g, groups[g].len(), w));
18047            groups[g].push(r);
18048        }
18049    }
18050    if experts.is_empty() {
18051        return out;
18052    }
18053    let pairs: Vec<_> = experts.iter().map(|e| (&e.gate_proj, &e.up_proj)).collect();
18054    let downs: Vec<_> = experts.iter().map(|e| &e.down_proj).collect();
18055    let lens: Vec<_> = groups.iter().map(Vec::len).collect();
18056    let count: usize = lens.iter().sum();
18057    let mut acts = vec![vec![0.0; experts[0].gate_proj.rows()]; count];
18058    let mut ds = vec![vec![0.0; hidden]; count];
18059    if QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut acts, pool)
18060        && QTensor::moe_down_rows(&downs, &lens, &acts, &mut ds, pool)
18061    {
18062        let mut offset = 0;
18063        let offsets: Vec<_> = lens
18064            .iter()
18065            .map(|&n| {
18066                let start = offset;
18067                offset += n;
18068                start
18069            })
18070            .collect();
18071        for (r, terms) in terms.iter().enumerate() {
18072            for &(g, slot, w) in terms {
18073                for (o, &v) in out[r * hidden..(r + 1) * hidden]
18074                    .iter_mut()
18075                    .zip(&ds[offsets[g] + slot])
18076                {
18077                    *o += w * v;
18078                }
18079            }
18080        }
18081    } else {
18082        for (r, jobs) in jobs.iter().enumerate() {
18083            if !jobs.is_empty() {
18084                let mut row = moe_cold_experts_cpu(jobs, &xs[r * hidden..(r + 1) * hidden], pool);
18085                out[r * hidden..(r + 1) * hidden].copy_from_slice(&row);
18086                attention::recycle_buf(&mut row);
18087            }
18088        }
18089    }
18090    out
18091}
18092
18093/// The pure-CPU MoE expert loop (also the fallback of every GPU refusal).
18094fn moe_ffn_cpu(
18095    m: &MoeFfn,
18096    x: &[f32],
18097    idx: &[usize],
18098    p: &[f32],
18099    wsum: f32,
18100    pool: Option<&Pool>,
18101) -> Vec<f32> {
18102    if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
18103        return out;
18104    }
18105    let mut out = attention::take_buf(x.len());
18106    for &e in idx {
18107        let mut eo = dense_ffn(&m.experts[e], x, pool);
18108        let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
18109        for i in 0..out.len() {
18110            out[i] += w * eo[i];
18111        }
18112        attention::recycle_buf(&mut eo);
18113    }
18114    if let Some((se, gate)) = &m.shared {
18115        let mut so = dense_ffn(se, x, pool);
18116        let g = gate.as_ref().map_or(1.0, |gate| {
18117            let mut gl = [0.0f32; 1];
18118            gate.matvec(x, &mut gl, pool);
18119            1.0 / (1.0 + (-gl[0]).exp())
18120        });
18121        for i in 0..out.len() {
18122            out[i] += g * so[i];
18123        }
18124        attention::recycle_buf(&mut so);
18125    }
18126    out
18127}
18128
18129/// DeepSeek-V2 MLA forward, expand-to-MHA form (see `AttnKind::Mla`):
18130/// per token the latent expands to every head's K/V and the ordinary
18131/// cache + grouped attend do the rest. K head layout is [rope | nope]
18132/// (rotary_dim = qk_rope rotates the shared rope key and each q head's
18133/// prefix); V rows are zero-padded to the K head_dim inside the cache
18134/// and the pad is sliced off before O. Attention importance is not
18135/// accumulated for MLA yet (no eviction interplay).
18136#[allow(clippy::too_many_arguments)]
18137pub(crate) fn mla_attention(
18138    w: &MlaWeights,
18139    normed: &[f32],
18140    cache: &mut crate::kv_cache::LayerKvCache,
18141    position: usize,
18142    inv_freq: &[f32],
18143    rope_scale: f32,
18144    eps: f64,
18145    pool: Option<&Pool>,
18146) -> Vec<f32> {
18147    let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
18148    let hd = dr + dn;
18149    let mut q = vec![0.0f32; nh * hd];
18150    match (&w.q_a, &w.q_a_norm) {
18151        (Some(qa), Some(qn)) => {
18152            let mut t = vec![0.0f32; qa.rows()];
18153            qa.matvec(normed, &mut t, pool);
18154            let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
18155            w.q_proj.matvec(&tn, &mut q, pool);
18156        }
18157        _ => w.q_proj.matvec(normed, &mut q, pool),
18158    }
18159    let mut ca = vec![0.0f32; lora + dr];
18160    w.kv_a.matvec(normed, &mut ca, pool);
18161    let (c_lat, k_rope) = ca.split_at_mut(lora);
18162    let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
18163    let mut kvb = vec![0.0f32; nh * (dn + dv)];
18164    w.kv_b.matvec(&latn, &mut kvb, pool);
18165    if !w.nope {
18166        attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
18167    }
18168    for h in 0..nh {
18169        if !w.nope {
18170            attention::rope_rotate_scaled(
18171                &mut q[h * hd..h * hd + dr],
18172                position,
18173                inv_freq,
18174                rope_scale,
18175            );
18176        }
18177    }
18178    let mut k = vec![0.0f32; nh * hd];
18179    let mut v = vec![0.0f32; nh * hd];
18180    for h in 0..nh {
18181        k[h * hd..h * hd + dr].copy_from_slice(k_rope);
18182        k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
18183        v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
18184    }
18185    cache.append(&k, &v, &vec![true; nh]);
18186    let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
18187    attention::recycle_buf(&mut imp);
18188    let mut ov = vec![0.0f32; nh * dv];
18189    for h in 0..nh {
18190        ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
18191    }
18192    let mut out = vec![0.0f32; w.o_proj.rows()];
18193    w.o_proj.matvec(&ov, &mut out, pool);
18194    out
18195}
18196
18197/// Gemma-4 dual-branch FFN (spec: see `FfnKind::DenseMoe`). The dense
18198/// branch reads the pre-FFN-normed activation; the router and the
18199/// expert branch read the RAW residual — the router through a
18200/// scale-less rms norm (its constant gain is folded into the weights),
18201/// the experts through `pre_norm_2`. CPU path; GPU graphs refuse the
18202/// layer kind honestly.
18203fn dense_moe_ffn(
18204    dm: &DenseMoeFfn,
18205    x_normed: &[f32],
18206    h_raw: &[f32],
18207    eps: f64,
18208    norm_style: NormStyle,
18209    pool: Option<&Pool>,
18210) -> Vec<f32> {
18211    let mut d = dense_ffn(&dm.dense, x_normed, pool);
18212    d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
18213    let m = &dm.moe;
18214    let ne = m.experts.len();
18215    let mut logits = vec![0.0f32; ne];
18216    if m.router_input_norm {
18217        let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
18218        let inv = 1.0 / (ss + eps as f32).sqrt();
18219        let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
18220        m.router.matvec(&xr, &mut logits, pool);
18221    } else {
18222        m.router.matvec(h_raw, &mut logits, pool);
18223    }
18224    let (idx, p, wsum) = moe_route(&logits, m, None);
18225    {
18226        let mut st = m.stats.borrow_mut();
18227        if st.len() < ne {
18228            st.resize(ne, 0);
18229        }
18230        for &e in &idx {
18231            st[e] += 1;
18232        }
18233    }
18234    let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
18235    let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
18236    let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
18237    for (di, mi) in d.iter_mut().zip(&mo) {
18238        *di += mi;
18239    }
18240    d
18241}
18242
18243/// Building the MoE-layer GPU jobs: all selected experts (+shared) must
18244/// be q8_2f-Mapped from the primary mapping; otherwise None → CPU path.
18245/// One-shot report of why the MoE GPU block refused. A silent `?` here
18246/// sends every expert to the CPU with nothing in the logs to say so —
18247/// which is exactly how a q4tp MoE model looked "GPU-accelerated" while
18248/// running entirely on the host.
18249fn moe_gpu_refused(why: &'static str) {
18250    use std::sync::atomic::{AtomicBool, Ordering};
18251    static SAID: AtomicBool = AtomicBool::new(false);
18252    if !SAID.swap(true, Ordering::Relaxed) {
18253        tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
18254    }
18255}
18256
18257fn moe_ffn_gpu(
18258    m: &MoeFfn,
18259    x: &[f32],
18260    idx: &[usize],
18261    p: &[f32],
18262    wsum: f32,
18263    pool: Option<&Pool>,
18264) -> Option<Vec<f32>> {
18265    use crate::gpu::MoeJob;
18266
18267    let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
18268    let mut model_ref = None;
18269    for &e in idx {
18270        if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
18271            moe_gpu_refused("push_job(expert)");
18272            return None;
18273        }
18274    }
18275    if let Some((se, gate)) = &m.shared {
18276        let g = gate.as_ref().map_or(1.0, |gate| {
18277            let mut gl = [0.0f32; 1];
18278            gate.matvec(x, &mut gl, pool);
18279            1.0 / (1.0 + (-gl[0]).exp())
18280        });
18281        if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
18282            moe_gpu_refused("push_job(shared)");
18283            return None;
18284        }
18285    }
18286    let Some(model) = model_ref else {
18287        moe_gpu_refused("no model_ref");
18288        return None;
18289    };
18290    let hidden = jobs[0].down.1;
18291    let mut out = vec![0.0f32; hidden];
18292    if crate::gpu::moe_block(&model, &jobs, &mut out) {
18293        Some(out)
18294    } else {
18295        moe_gpu_refused("gpu::moe_block");
18296        None
18297    }
18298}
18299
18300/// Single-position FFN dispatch.
18301fn ffn_forward(
18302    ffn: &FfnKind,
18303    x: &[f32],
18304    pool: Option<&Pool>,
18305    experts_allowed: Option<&[bool]>,
18306) -> Vec<f32> {
18307    match ffn {
18308        FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
18309        FfnKind::Dense(d) => dense_ffn(d, x, pool),
18310        FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
18311        // Dual-branch layers need the raw residual — their callers
18312        // dispatch dense_moe_ffn directly; the auxiliary paths that land
18313        // here (MTP draft, o1 replay) do not co-occur with gemma-4 MoE.
18314        FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18315    }
18316}
18317
18318/// Fused two-position FFN: gate/up/down streamed once (dense). MoE
18319/// falls back to two singles — expert sets differ per position, there
18320/// is nothing to fuse.
18321fn ffn_forward_pair(
18322    ffn: &FfnKind,
18323    x1: &[f32],
18324    x2: &[f32],
18325    pool: Option<&Pool>,
18326    experts_allowed: Option<&[bool]>,
18327) -> (Vec<f32>, Vec<f32>) {
18328    let d = match ffn {
18329        // A tube layer has nothing to fuse across the pair — the tubes
18330        // are separate matrices; two singles are the honest path.
18331        FfnKind::Dense(d) if !d.segs.is_empty() => {
18332            return (
18333                tube_ffn(d, x1, 1, pool, None),
18334                tube_ffn(d, x2, 1, pool, None),
18335            );
18336        }
18337        FfnKind::Dense(d) => d,
18338        FfnKind::Moe(m) => {
18339            return (
18340                moe_ffn(m, x1, pool, experts_allowed),
18341                moe_ffn(m, x2, pool, experts_allowed),
18342            );
18343        }
18344        FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18345    };
18346    let inter = d.gate_proj.rows();
18347    FFN_SCRATCH.with(|s| {
18348        let mut s = s.borrow_mut();
18349        let [g1, g2, u1, u2] = &mut *s;
18350        g1.resize(inter, 0.0);
18351        g2.resize(inter, 0.0);
18352        u1.resize(inter, 0.0);
18353        u2.resize(inter, 0.0);
18354        // Multi-matrix pair job: gate+up under one pool dispatch
18355        // (o1s = lane-1 outputs across tensors, o2s = lane-2).
18356        QTensor::matvec2_many(
18357            [&d.gate_proj, &d.up_proj],
18358            x1,
18359            x2,
18360            [g1.as_mut_slice(), u1.as_mut_slice()],
18361            [g2.as_mut_slice(), u2.as_mut_slice()],
18362            pool,
18363        );
18364        for i in 0..inter {
18365            g1[i] = d.act.combine(g1[i], u1[i]);
18366            g2[i] = d.act.combine(g2[i], u2[i]);
18367        }
18368        let mut o1 = attention::take_buf(d.down_proj.rows());
18369        let mut o2 = attention::take_buf(d.down_proj.rows());
18370        d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
18371        (o1, o2)
18372    })
18373}
18374
18375#[cfg(test)]
18376mod tests {
18377
18378    /// The 0.7.6 prefill-chunk rule: a plain dense stack wholly on a
18379    /// discrete card reads the prompt in wide chunks on x86; every other
18380    /// case keeps the width it had (the GDN-hybrid, MoE and DeepSeek paths
18381    /// were tuned on hardware not measured for this change).
18382    #[test]
18383    fn prefill_chunk_rule_widens_only_dense_on_discrete() {
18384        use super::{
18385            prefill_chunk_rule, ChunkHost, ChunkStackFacts, DISCRETE_DENSE_PREFILL_CHUNK,
18386        };
18387        let dense_card = ChunkStackFacts {
18388            plain_dense: true,
18389            discrete: true,
18390            gpu_on: true,
18391            ..Default::default()
18392        };
18393        assert!(dense_card.dense_on_discrete());
18394        // The bug: a dense Llama on a Vulkan RTX 3090 got 48.
18395        assert_eq!(
18396            prefill_chunk_rule(None, ChunkHost::Other, dense_card.dense_on_discrete()),
18397            DISCRETE_DENSE_PREFILL_CHUNK
18398        );
18399        assert!(DISCRETE_DENSE_PREFILL_CHUNK > 48);
18400        for (label, facts) in [
18401            ("GDN hybrid / MoE / DeepSeek stack", ChunkStackFacts { plain_dense: false, ..dense_card }),
18402            ("integrated GPU", ChunkStackFacts { discrete: false, ..dense_card }),
18403            ("CPU only", ChunkStackFacts { gpu_on: false, discrete: false, ..dense_card }),
18404            ("capacity split", ChunkStackFacts { capacity_split: true, ..dense_card }),
18405            ("multi-GPU plan", ChunkStackFacts { multi_gpu: true, ..dense_card }),
18406            ("O(1) layers", ChunkStackFacts { o1: true, ..dense_card }),
18407        ] {
18408            assert!(!facts.dense_on_discrete(), "{label}");
18409            assert_eq!(
18410                prefill_chunk_rule(None, ChunkHost::Other, facts.dense_on_discrete()),
18411                48,
18412                "{label} keeps the historical x86 chunk"
18413            );
18414        }
18415        // Other hosts are untouched whatever the model.
18416        for dense in [false, true] {
18417            assert_eq!(prefill_chunk_rule(None, ChunkHost::Macos, dense), 512);
18418            assert_eq!(prefill_chunk_rule(None, ChunkHost::Aarch64, dense), 256);
18419        }
18420        // CMF_PREFILL_CHUNK still wins everywhere (and is clamped to ≥ 1).
18421        for host in [ChunkHost::Macos, ChunkHost::Aarch64, ChunkHost::Other] {
18422            for dense in [false, true] {
18423                assert_eq!(prefill_chunk_rule(Some(48), host, dense), 48);
18424                assert_eq!(prefill_chunk_rule(Some(0), host, dense), 1);
18425            }
18426        }
18427    }
18428
18429    #[test]
18430    fn kv_reuse_plan_pulls_rows_decode_wrote_only_on_the_device() {
18431        use super::{ReuseLayer, ReusePlan, kv_reuse_plan};
18432        let full = |host_rows, device_rows| ReuseLayer {
18433            full: true,
18434            host_rows,
18435            device_rows,
18436            device_state: false,
18437        };
18438        // Turn 1: 300-token prompt prefilled on the host, 40 tokens decoded
18439        // by the wgpu graph into the device mirror only. Turn 2 reuses 339.
18440        assert_eq!(
18441            kv_reuse_plan(339, &[full(300, Some(339)), full(300, Some(339))]),
18442            ReusePlan::Pull(vec![(0, 300, 339), (1, 300, 339)])
18443        );
18444        // CPU / Metal: the host owner already holds every forwarded row.
18445        assert_eq!(kv_reuse_plan(339, &[full(339, None)]), ReusePlan::Ready);
18446        // A mirror past the prefix is fine for the host (it gets rewound).
18447        assert_eq!(kv_reuse_plan(339, &[full(339, Some(345))]), ReusePlan::Ready);
18448        // GPU prefix / CPU tail: only the device layers lag.
18449        assert_eq!(
18450            kv_reuse_plan(339, &[full(300, Some(339)), full(339, None)]),
18451            ReusePlan::Pull(vec![(0, 300, 339)])
18452        );
18453        // The device cannot supply the missing rows: never continue.
18454        assert_eq!(kv_reuse_plan(339, &[full(300, Some(320))]), ReusePlan::Fresh);
18455        assert_eq!(kv_reuse_plan(339, &[full(300, None)]), ReusePlan::Fresh);
18456        assert_eq!(kv_reuse_plan(339, &[full(350, None)]), ReusePlan::Fresh);
18457        // A recurrent state advanced on the device cannot be handed to a
18458        // host prefill (it is not rewindable and the host copy is stale).
18459        let conv = |device_state| ReuseLayer {
18460            full: false,
18461            host_rows: 0,
18462            device_rows: None,
18463            device_state,
18464        };
18465        assert_eq!(
18466            kv_reuse_plan(339, &[conv(true), full(300, Some(339))]),
18467            ReusePlan::Fresh
18468        );
18469        assert_eq!(kv_reuse_plan(339, &[conv(false), full(339, None)]), ReusePlan::Ready);
18470    }
18471
18472    #[test]
18473    fn nll_graph_policy_scopes_only_the_fused_head() {
18474        for (label, unmasked, prefer_graph, native_metal, want_graph, want_head) in [
18475            // A Vulkan/Wgpu hidden-only graph remains the quality route.
18476            ("vulkan graph", true, true, false, true, false),
18477            // Native Metal adds the strict fused graph-head contract.
18478            ("native Metal graph", true, true, true, true, true),
18479            // Masked NLL and the explicit non-graph fallback remain unchanged.
18480            ("masked", false, true, false, false, false),
18481            ("graph disabled", true, false, true, false, false),
18482        ] {
18483            let (graph_quality, graph_head_required) =
18484                super::nll_graph_policy(unmasked, prefer_graph, native_metal);
18485            assert_eq!(graph_quality, want_graph, "{label}: graph quality");
18486            assert_eq!(graph_head_required, want_head, "{label}: fused head");
18487        }
18488    }
18489
18490    #[test]
18491    fn mtp_prefill_pair_boundaries_skip_only_final_prompt_row() {
18492        assert_eq!(mtp_prefill_pair_count(0, 128, 256), 128);
18493        assert_eq!(mtp_prefill_pair_count(128, 256, 256), 127);
18494        assert_eq!(mtp_prefill_pair_count(0, 256, 256), 255);
18495        assert_eq!(mtp_prefill_pair_count(256, 256, 256), 0);
18496        assert_eq!(mtp_prefill_pair_count(300, 320, 256), 0);
18497    }
18498
18499    #[test]
18500    fn cancel_flag_stops_generation() {
18501        let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
18502        // Set before the call: the prefill loops honour it, the run
18503        // returns immediately with the cancelled reason and no tokens.
18504        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
18505        let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
18506        assert_eq!(r.finish_reason, "cancelled");
18507        assert!(
18508            r.token_ids.is_empty(),
18509            "no tokens after cancel: {:?}",
18510            r.token_ids
18511        );
18512        assert_eq!(p.kv_cache.seq_len(), 0);
18513        assert!(p.kv_history.is_empty());
18514        assert!(!p.graph_want_logits);
18515        assert!(p.graph_logits.is_none());
18516        // Flag auto-cleared: the next call generates normally.
18517        let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
18518        assert_ne!(r2.finish_reason, "cancelled");
18519    }
18520    use super::*;
18521
18522    /// sparse_ffn_quant must equal a dense FFN where inactive neurons are
18523    /// zeroed (mask × mmap correctness). On F32 tensors this is EXACT —
18524    /// it validates the row_dot / add_col_scaled / scatter indexing, the
18525    /// bug-prone part. The q8 branches reuse the golden-tested linear
18526    /// The per-token sparse path reads a transposed `down`; it must
18527    /// agree with the arm that computes everything and zeroes the
18528    /// losers, or the speed measurement is measuring a different model.
18529    #[test]
18530    fn dynamic_ffn_equals_the_zeroing_arm() {
18531        let (hidden, inter) = (8usize, 32usize);
18532        let synth = |n: usize, salt: usize| -> Vec<f32> {
18533            (0..n)
18534                .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
18535                .collect()
18536        };
18537        let down = synth(hidden * inter, 3);
18538        let mut down_t = vec![0.0f32; inter * hidden];
18539        for r in 0..hidden {
18540            for c in 0..inter {
18541                down_t[c * hidden + r] = down[r * inter + c];
18542            }
18543        }
18544        let d = DenseFfn {
18545            gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18546            up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18547            down_proj: QTensor::from_f32(down.clone(), hidden, inter),
18548            act: Act::Silu,
18549            down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
18550            segs: Vec::new(),
18551        };
18552        let x = synth(hidden, 11);
18553        let k = 12usize;
18554        let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
18555        // Reference: full compute, keep the k loudest |silu(gate)|.
18556        let mut g = vec![0.0f32; inter];
18557        d.gate_proj.matvec(&x, &mut g, None);
18558        let mut u = vec![0.0f32; inter];
18559        d.up_proj.matvec(&x, &mut u, None);
18560        for v in g.iter_mut() {
18561            *v = inference::silu(*v);
18562        }
18563        keep_top_k(&mut g, k);
18564        for i in 0..inter {
18565            g[i] *= u[i];
18566        }
18567        let mut want = vec![0.0f32; hidden];
18568        d.down_proj.matvec(&g, &mut want, None);
18569        for (a, b) in want.iter().zip(&got) {
18570            assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
18571        }
18572    }
18573
18574    /// A tube layer is the same layer, re-cut. With every tube open the
18575    /// answer must equal the dense FFN over the concatenated neurons
18576    /// (the permutation is an identity on the layer's function); with a
18577    /// tube closed it must equal the dense FFN with those neurons
18578    /// zeroed — the mask semantics, now paid for in bytes not read.
18579    #[test]
18580    fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
18581        let (hidden, core, tube) = (8usize, 12usize, 8usize);
18582        let inter = core + tube;
18583        let synth = |n: usize, salt: usize| -> Vec<f32> {
18584            (0..n)
18585                .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
18586                .collect()
18587        };
18588        let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
18589        let d_all = synth(hidden * inter, 3);
18590        // The dense layer, and the same weights cut into core + tube.
18591        let dense = DenseFfn {
18592            gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
18593            up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
18594            down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
18595            act: Act::Silu,
18596            down_t: None,
18597            segs: Vec::new(),
18598        };
18599        let rows =
18600            |v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
18601        let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
18602            let mut o = Vec::with_capacity(hidden * (b - a));
18603            for r in 0..hidden {
18604                o.extend_from_slice(&v[r * inter + a..r * inter + b]);
18605            }
18606            o
18607        };
18608        let tubed = DenseFfn {
18609            down_t: None,
18610            gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
18611            up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
18612            down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
18613            act: Act::Silu,
18614            segs: vec![FfnSeg {
18615                gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
18616                up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
18617                down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
18618                start: core,
18619                width: tube,
18620            }],
18621        };
18622        let x = synth(hidden, 7);
18623        let want = dense_ffn(&dense, &x, None);
18624        let got = tube_ffn(&tubed, &x, 1, None, None);
18625        for (a, b) in want.iter().zip(&got) {
18626            assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
18627        }
18628        // Closed tube: bits on for the core, off for the tube.
18629        let mut bits = vec![0u8; inter.div_ceil(8)];
18630        for n in 0..core {
18631            bits[n / 8] |= 1 << (n % 8);
18632        }
18633        let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18634        let masked = dense_ffn_masked(&dense, &x, None, &bits);
18635        for (a, b) in masked.iter().zip(&closed) {
18636            assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
18637        }
18638        // The batched arm must agree with the single-position one.
18639        let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18640        for (a, b) in closed.iter().zip(&batch) {
18641            assert_eq!(a, b, "batch arm disagrees with decode arm");
18642        }
18643    }
18644
18645    /// scale, structurally identical to the matvec kernels.
18646    #[test]
18647    fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
18648        let (hidden, inter) = (16usize, 40usize);
18649        let synth = |n: usize, salt: usize| -> Vec<f32> {
18650            (0..n)
18651                .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
18652                .collect()
18653        };
18654        let d = DenseFfn {
18655            gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18656            up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18657            down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
18658            act: Act::Silu,
18659            down_t: None,
18660            segs: Vec::new(),
18661        };
18662        let x = synth(hidden, 9);
18663        // Active = every 3rd neuron.
18664        let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
18665
18666        let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
18667
18668        // Reference: full dense FFN but g[i]=0 for inactive neurons.
18669        let mut g = vec![0.0f32; inter];
18670        d.gate_proj.matvec(&x, &mut g, None);
18671        let mut u = vec![0.0f32; inter];
18672        d.up_proj.matvec(&x, &mut u, None);
18673        let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
18674        for i in 0..inter {
18675            g[i] = if act_set.contains(&(i as u16)) {
18676                inference::silu(g[i]) * u[i]
18677            } else {
18678                0.0
18679            };
18680        }
18681        let mut reference = vec![0.0f32; hidden];
18682        d.down_proj.matvec(&g, &mut reference, None);
18683
18684        let max_d = sparse
18685            .iter()
18686            .zip(&reference)
18687            .map(|(a, b)| (a - b).abs())
18688            .fold(0.0f32, f32::max);
18689        assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
18690    }
18691
18692    /// Attach a synthetic MTP head (same structure as a main layer).
18693    fn attach_test_mtp(p: &mut Pipeline) {
18694        let (h, inter, heads, kv, hd) = (
18695            p.hidden_size,
18696            p.intermediate_size,
18697            p.num_heads,
18698            p.num_kv_heads,
18699            p.head_dim,
18700        );
18701        let synth = |n: usize, salt: usize| -> Vec<f32> {
18702            (0..n)
18703                .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
18704                .collect()
18705        };
18706        let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
18707            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
18708        };
18709        p.mtp = Some(MtpModule {
18710            enorm: vec![1.0; h],
18711            hnorm: vec![1.0; h],
18712            eh_proj: qt(h, 2 * h, 301),
18713            layer: LayerWeights {
18714                input_norm: vec![1.0; h],
18715                post_norm: vec![1.0; h],
18716                attn_out_norm: None,
18717                ffn_out_norm: None,
18718                layer_scale: None,
18719                ffn: FfnKind::Dense(DenseFfn {
18720                    gate_proj: qt(inter, h, 315),
18721                    up_proj: qt(inter, h, 316),
18722                    down_proj: qt(h, inter, 317),
18723                    act: Act::Silu,
18724                    down_t: None,
18725                    segs: Vec::new(),
18726                }),
18727                attn: AttnKind::Full {
18728                    bias: None,
18729                    wq: qt(heads * hd, h, 311),
18730                    wk: qt(kv * hd, h, 312),
18731                    wv: qt(kv * hd, h, 313),
18732                    wo: qt(h, heads * hd, 314),
18733                    q_norm: None,
18734                    k_norm: None,
18735                    output_gate: false,
18736                    softplus_gate: None,
18737                },
18738            },
18739            final_norm: vec![1.0; h],
18740            kv: crate::kv_cache::LayerKvCache::new(kv, hd),
18741        });
18742    }
18743
18744    #[test]
18745    fn speculative_equals_vanilla_greedy() {
18746        // Speculative decode and the wgpu token graph are mutually
18747        // exclusive; a leaked CMF_GPU=wgpu from a parallel gpu test
18748        // would silently disable drafting. Pin the graph off.
18749        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
18750        let run = |spec: bool| {
18751            let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18752            p.sampler_config.temperature = 0.0;
18753            attach_test_mtp(&mut p);
18754            p.speculative = spec;
18755            let r = p.generate("abcdef", 12, None, None).unwrap();
18756            (r.token_ids, r.mtp_drafted, r.mtp_accepted)
18757        };
18758        let (vanilla, d0, _) = run(false);
18759        let (spec, d1, a1) = run(true);
18760        assert_eq!(d0, 0, "vanilla path must not draft");
18761        assert!(d1 > 0, "speculative path must draft");
18762        assert_eq!(
18763            vanilla, spec,
18764            "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
18765        );
18766    }
18767
18768    #[test]
18769    fn speculative_accepts_constant_oracle() {
18770        // See speculative_equals_vanilla_greedy: pin the wgpu graph off.
18771        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
18772        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18773        p.sampler_config.temperature = 0.0;
18774        p.sampler_config.repetition_penalty = 1.0;
18775        // Constant lm_head → every logit equal → both the main model and
18776        // the draft head argmax to token 0: acceptance must be 100%.
18777        p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
18778        attach_test_mtp(&mut p);
18779        p.speculative = true;
18780        let r = p.generate("abcd", 10, None, None).unwrap();
18781        assert!(r.mtp_drafted > 0);
18782        assert_eq!(
18783            r.mtp_accepted, r.mtp_drafted,
18784            "constant logits → every draft accepted"
18785        );
18786        // Ties resolve to the same token in both the main and draft
18787        // heads — the sequence is one repeated token.
18788        assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
18789    }
18790
18791    #[test]
18792    fn empty_prompt_is_an_error_not_a_panic() {
18793        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
18794        let r = p.generate("", 4, None, None);
18795        assert!(r.is_err(), "empty prompt must be a clean error");
18796    }
18797
18798    #[test]
18799    fn every_token_enters_kv_exactly_once() {
18800        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18801        // Greedy so no RNG variance; byte tokenizer → 3 prompt tokens.
18802        p.sampler_config.temperature = 0.0;
18803        let r = p.generate("abc", 2, None, None).unwrap();
18804        assert_eq!(r.prompt_tokens, 3);
18805        // prompt(3) + first sampled token forwarded before second logits:
18806        // step0 samples from prefill hidden (no extra forward), then
18807        // forwards t1 → cache 4; step1 samples, loop ends (max_tokens).
18808        assert_eq!(
18809            p.kv_cache.seq_len(),
18810            3 + r.tokens_generated - 1,
18811            "each token must be cached exactly once (v1 cached the last prompt token twice)"
18812        );
18813    }
18814
18815    #[test]
18816    fn generation_is_reproducible_with_seed() {
18817        let run = || {
18818            let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18819            p.generate("hello", 8, None, None).unwrap().token_ids
18820        };
18821        assert_eq!(run(), run());
18822    }
18823
18824    #[test]
18825    fn resetting_sampler_restarts_the_seeded_stream() {
18826        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18827        let config = SamplerConfig {
18828            seed: Some(1234),
18829            ..SamplerConfig::default()
18830        };
18831        p.set_sampler_config(config.clone());
18832        let first = p.generate("hello", 8, None, None).unwrap().token_ids;
18833        p.set_sampler_config(config);
18834        let second = p.generate("hello", 8, None, None).unwrap().token_ids;
18835        assert_eq!(first, second);
18836    }
18837
18838    #[test]
18839    fn eviction_bounds_the_cache() {
18840        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
18841        p.kv_cache.max_seq_len = 6;
18842        p.sampler_config.temperature = 0.0;
18843        let _ = p.generate("abcd", 12, None, None).unwrap();
18844        assert!(
18845            p.kv_cache.seq_len() <= 6 + 1,
18846            "cache must stay bounded by max_seq_len (got {})",
18847            p.kv_cache.seq_len()
18848        );
18849    }
18850
18851    #[test]
18852    fn confidence_matches_tokens_and_is_a_probability() {
18853        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18854        p.sampler_config.temperature = 0.0;
18855        p.sampler_config.repetition_penalty = 1.0;
18856        let r = p.generate("abcd", 10, None, None).unwrap();
18857        assert_eq!(
18858            r.token_confidence.len(),
18859            r.token_ids.len(),
18860            "one confidence per emitted token"
18861        );
18862        for &c in &r.token_confidence {
18863            assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
18864        }
18865        // top1_prob is a valid softmax probability.
18866        let logits = [1.0f32, 3.0, 0.5, 3.0];
18867        let p0 = top1_prob_t(&logits, 1, 1.0);
18868        let p1 = top1_prob_t(&logits, 3, 1.0);
18869        assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
18870        assert!(p0 > 0.0 && p0 < 1.0);
18871        // Calibration temperature > 1 softens an over-confident peak.
18872        let sharp = top1_prob_t(&logits, 1, 1.0);
18873        let soft = top1_prob_t(&logits, 1, 2.0);
18874        assert!(soft < sharp, "higher temperature lowers peak confidence");
18875    }
18876
18877    #[test]
18878    fn trace_is_opt_in_and_parallels_the_output() {
18879        // Off by default: the runtime is silent unless observation asked.
18880        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18881        p.sampler_config.temperature = 0.0;
18882        p.sampler_config.repetition_penalty = 1.0;
18883        let r = p.generate("abcd", 10, None, None).unwrap();
18884        assert!(r.traces.is_empty(), "trace must be empty unless enabled");
18885
18886        // On: exactly one row per emitted token, aligned with the output.
18887        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18888        p.sampler_config.temperature = 0.0;
18889        p.sampler_config.repetition_penalty = 1.0;
18890        p.set_trace(true);
18891        let r = p.generate("abcd", 10, None, None).unwrap();
18892        assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
18893        for (i, tr) in r.traces.iter().enumerate() {
18894            assert_eq!(tr.t, i, "trace index is sequential");
18895            assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
18896            assert_eq!(
18897                tr.confidence, r.token_confidence[i],
18898                "trace confidence matches the confidence channel"
18899            );
18900            // No dynamic router in this pipeline → no skill, no coherence.
18901            assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
18902        }
18903    }
18904
18905    #[test]
18906    fn explain_prefill_logits_match_greedy_first_token() {
18907        // `cortiq explain` shows the next-token distribution from
18908        // prefill_next_logits; its argmax must equal what greedy generate
18909        // actually emits first — otherwise explain would lie.
18910        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18911        p.sampler_config.temperature = 0.0;
18912        p.sampler_config.repetition_penalty = 1.0;
18913        let ids = p.tokenizer.encode("abcd");
18914        let logits = p.prefill_next_logits(&ids, None);
18915        let argmax = logits
18916            .iter()
18917            .enumerate()
18918            .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
18919            .unwrap()
18920            .0 as u32;
18921        let r = p.generate("abcd", 1, None, None).unwrap();
18922        assert_eq!(
18923            argmax, r.token_ids[0],
18924            "explain preview must match greedy emit"
18925        );
18926    }
18927
18928    #[test]
18929    fn laguna_shared_expert_is_unconditionally_added() {
18930        let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
18931        let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
18932        let zero_dense = || DenseFfn {
18933            gate_proj: matrix(vec![0.0; 4]),
18934            up_proj: matrix(vec![0.0; 4]),
18935            down_proj: matrix(vec![0.0; 4]),
18936            act: Act::Silu,
18937            down_t: None,
18938            segs: Vec::new(),
18939        };
18940        let shared = DenseFfn {
18941            gate_proj: identity(),
18942            up_proj: identity(),
18943            down_proj: identity(),
18944            act: Act::Silu,
18945            down_t: None,
18946            segs: Vec::new(),
18947        };
18948        let x = [1.0, 2.0];
18949        let expected = dense_ffn(&shared, &x, None);
18950        let moe = MoeFfn {
18951            router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
18952            experts: vec![zero_dense()],
18953            top_k: 1,
18954            norm_topk_prob: true,
18955            router_sigmoid: true,
18956            expert_bias: None,
18957            routed_scaling: 1.0,
18958            route_tau: None,
18959            shared: Some((shared, None)),
18960            stats: std::cell::RefCell::new(Vec::new()),
18961            act_sq: std::cell::RefCell::new(Vec::new()),
18962            act_rows: std::cell::RefCell::new(Vec::new()),
18963            mask: None,
18964            per_expert_scale: None,
18965            router_input_norm: false,
18966            resonance: None,
18967            grown: Vec::new(),
18968        };
18969        let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
18970        for (actual, expected) in actual.iter().zip(expected) {
18971            assert!((actual - expected).abs() < 1e-6);
18972        }
18973    }
18974
18975    /// A tiny MiMo-V2-shaped stack (the M3 fixture): layers [full, sliding,
18976    /// sliding, full]; 4 Q heads over 1 (full) / 2 (sliding) KV heads;
18977    /// head_dim 8 with 4-wide V heads; partial rotary 4 at θ 1e7 (full) /
18978    /// 1e4 (sliding); window 3; learned sinks on the sliding layers; layer
18979    /// 0 a dense FFN, layers 1..3 sigmoid-routed MoE with a selection bias
18980    /// (4 experts, top-2, renormalized, no shared expert). Geometry and
18981    /// sinks go through the same `set_attn_geometry` / `set_layer_sinks`
18982    /// the loader calls.
18983    fn mimo_test_pipeline() -> Pipeline {
18984        let (hs, inter, nh, hd, vd, vocab) = (16usize, 24usize, 4usize, 8usize, 4usize, 64usize);
18985        let kvh = [1usize, 2, 2, 1];
18986        let synth = |n: usize, salt: usize| -> Vec<f32> {
18987            (0..n)
18988                .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
18989                .collect()
18990        };
18991        let qt = |rows: usize, cols: usize, salt: usize| {
18992            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
18993        };
18994        let dense = |inter: usize, salt: usize| DenseFfn {
18995            gate_proj: qt(inter, hs, salt),
18996            up_proj: qt(inter, hs, salt + 1),
18997            down_proj: qt(hs, inter, salt + 2),
18998            act: Act::Silu,
18999            down_t: None,
19000            segs: Vec::new(),
19001        };
19002        let layers: Vec<LayerWeights> = (0..4)
19003            .map(|li| LayerWeights {
19004                input_norm: vec![1.0; hs],
19005                post_norm: vec![1.0; hs],
19006                attn_out_norm: None,
19007                ffn_out_norm: None,
19008                layer_scale: None,
19009                ffn: if li == 0 {
19010                    FfnKind::Dense(dense(inter, 50))
19011                } else {
19012                    FfnKind::Moe(MoeFfn {
19013                        router: qt(4, hs, 60 + li),
19014                        experts: (0..4).map(|e| dense(8, 70 + li * 10 + e * 3)).collect(),
19015                        top_k: 2,
19016                        norm_topk_prob: true,
19017                        router_sigmoid: true,
19018                        expert_bias: Some(vec![0.02, -0.03, 0.01, 0.0]),
19019                        routed_scaling: 1.0,
19020                        route_tau: None,
19021                        shared: None,
19022                        stats: std::cell::RefCell::new(Vec::new()),
19023                        act_sq: std::cell::RefCell::new(Vec::new()),
19024                        act_rows: std::cell::RefCell::new(Vec::new()),
19025                        mask: None,
19026                        per_expert_scale: None,
19027                        router_input_norm: false,
19028                        resonance: None,
19029                        grown: Vec::new(),
19030                    })
19031                },
19032                attn: AttnKind::Full {
19033                    wq: qt(nh * hd, hs, li * 10 + 1),
19034                    wk: qt(kvh[li] * hd, hs, li * 10 + 2),
19035                    wv: qt(kvh[li] * vd, hs, li * 10 + 3),
19036                    wo: qt(hs, nh * vd, li * 10 + 4),
19037                    q_norm: None,
19038                    k_norm: None,
19039                    output_gate: false,
19040                    softplus_gate: None,
19041                    bias: None,
19042                },
19043            })
19044            .collect();
19045        let mut p = Pipeline::new(
19046            Tokenizer::byte_level(),
19047            PipelineWeights {
19048                embed_tokens: qt(vocab, hs, 100),
19049                layers,
19050                lm_head: qt(vocab, hs, 200),
19051                final_norm: vec![1.0; hs],
19052            },
19053            hs,
19054            inter,
19055            nh,
19056            1, // header num_kv_heads (the full layers')
19057            hd,
19058            4,
19059            4,
19060            false,
19061            vocab,
19062            1e-6,
19063            1e7,
19064            NormStyle::Qwen,
19065            4096,
19066            SamplerConfig {
19067                seed: Some(7),
19068                ..Default::default()
19069            },
19070        );
19071        // Diagnostics stay off whatever the test environment exports.
19072        p.layer_dump = None;
19073        p.set_rotary(4, 1e7);
19074        p.sliding_layers = Some(vec![false, true, true, false]);
19075        p.swa = Some((3, usize::MAX));
19076        p.rotary_dim_local = Some(4);
19077        p.inv_freq_local = Some(std::sync::Arc::new(attention::rope_inv_freq(4, 1e4)));
19078        p.set_attn_geometry(Some(kvh.to_vec()), Some(vd)).unwrap();
19079        p.set_layer_sinks(1, vec![0.5, -1.0, 1.5, 0.0]).unwrap();
19080        p.set_layer_sinks(2, vec![-0.25, 2.0, 0.75, -1.5]).unwrap();
19081        p
19082    }
19083
19084    #[test]
19085    fn mimo_embedded_prompt_uses_rows_and_never_reuses_token_only_kv() {
19086        let mut p = mimo_test_pipeline();
19087        p.speculative = false;
19088        p.ignore_eos = true;
19089        p.sampler_config.temperature = 0.0;
19090        p.sampler_config.repetition_penalty = 1.0;
19091        let a = vec![3, 5, 7, 9, 11, 13];
19092        let b = vec![4, 8, 12, 16, 20, 24];
19093        let rows: Vec<_> = b.iter().flat_map(|&id| p.embed_id(id)).collect();
19094        let expected = p.generate_from_ids(&b, 8, None, None).unwrap().token_ids;
19095        // Same placeholder IDs as an earlier request are not a cache key
19096        // for different media. The actual rows, not a re-embedding of a,
19097        // must determine the continuation.
19098        let actual = p.generate_from_embeds(&a, &rows, 8, None, None).unwrap().token_ids;
19099        assert_eq!(actual, expected);
19100        assert!(p.kv_history.is_empty());
19101        let mut extended = a.clone();
19102        extended.push(17);
19103        let after_media = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19104        p.reset_session();
19105        let fresh = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19106        assert_eq!(after_media, fresh);
19107        assert!(p.generate_from_embeds(&a, &rows[..rows.len()-1], 1, None, None).is_err());
19108        // Force a real token-prefix reuse opportunity into the media call.
19109        // Those labels are unchanged, but their embeddings now describe a
19110        // different source sequence and every KV row must be rebuilt.
19111        p.reset_session();
19112        p.generate_from_ids(&a, 1, None, None).unwrap();
19113        let mut media_ids = p.kv_history.clone();
19114        assert!(!media_ids.is_empty());
19115        media_ids.extend_from_slice(&[19, 21, 23]);
19116        let source_ids: Vec<_> = (0..media_ids.len()).map(|i| b[i % b.len()]).collect();
19117        let source_rows: Vec<_> = source_ids.iter().flat_map(|&id| p.embed_id(id)).collect();
19118        let mut oracle = mimo_test_pipeline();
19119        oracle.speculative = false;
19120        oracle.ignore_eos = true;
19121        oracle.sampler_config.temperature = 0.0;
19122        oracle.sampler_config.repetition_penalty = 1.0;
19123        let expected = oracle.generate_from_ids(&source_ids, 8, None, None).unwrap().token_ids;
19124        assert_eq!(p.generate_from_embeds(&media_ids, &source_rows, 8, None, None).unwrap().token_ids, expected);
19125        assert!(p.kv_history.is_empty());
19126        let mut bad = rows;
19127        bad[0] = f32::NAN;
19128        assert!(p.generate_from_embeds(&a, &bad, 1, None, None).is_err());
19129    }
19130
19131    fn f32_bits(v: &[f32]) -> Vec<u32> {
19132        v.iter().map(|x| x.to_bits()).collect()
19133    }
19134
19135    /// M3 acceptance: on the MiMo-shaped stack the decode walk (one
19136    /// position at a time through `forward_layers`) and the batched
19137    /// prefill (`prefill_batch_span`, whole prompt and split in two
19138    /// chunks) give bit-identical logits at all 12 positions — per-layer
19139    /// KV heads, narrow V, sinks, the window and the biased sigmoid MoE all
19140    /// agree across the two walks.
19141    #[test]
19142    fn mimo_shaped_decode_matches_prefill_batch_bitwise() {
19143        let mut p = mimo_test_pipeline();
19144        let kv: Vec<usize> = p.kv_cache.layers.iter().map(|l| l.num_kv_heads).collect();
19145        assert_eq!(kv, vec![1, 2, 2, 1]);
19146        assert!(p.kv_cache.layers[0].sinks.is_none() && p.kv_cache.layers[3].sinks.is_none());
19147        assert!(p.kv_cache.layers[1].sinks.is_some() && p.kv_cache.layers[2].sinks.is_some());
19148        let ids: Vec<u32> = (0..12u32).map(|i| (i * 7 + 3) % 64).collect();
19149        let hs = p.hidden_size;
19150        let mut decode = Vec::new();
19151        for (pos, &id) in ids.iter().enumerate() {
19152            let e = p.embed_single(id);
19153            let h = p.forward_layers(&e, pos, None);
19154            decode.push(p.logits_from_hidden(&h));
19155        }
19156        for l in &p.kv_cache.layers {
19157            assert_eq!(l.seq_len, 12);
19158            // V rows are padded to head_dim inside the cache.
19159            assert_eq!(l.head_values(0).len(), 12 * 8);
19160        }
19161        assert!(decode.iter().flatten().all(|v| v.is_finite()));
19162
19163        p.clear_sequence_state();
19164        let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19165        for pos in 0..ids.len() {
19166            let lg = p.logits_from_hidden(&hb[pos * hs..(pos + 1) * hs]);
19167            assert_eq!(
19168                f32_bits(&decode[pos]),
19169                f32_bits(&lg),
19170                "whole prompt, pos {pos}"
19171            );
19172        }
19173
19174        p.clear_sequence_state();
19175        let a = p.prefill_batch_span(PrefillIn::Ids(&ids[..5]), 0, None, 0, p.num_layers);
19176        let b = p.prefill_batch_span(PrefillIn::Ids(&ids[5..]), 5, None, 0, p.num_layers);
19177        for pos in 0..ids.len() {
19178            let row = if pos < 5 {
19179                &a[pos * hs..(pos + 1) * hs]
19180            } else {
19181                &b[(pos - 5) * hs..(pos - 4) * hs]
19182            };
19183            let lg = p.logits_from_hidden(row);
19184            assert_eq!(
19185                f32_bits(&decode[pos]),
19186                f32_bits(&lg),
19187                "two chunks, pos {pos}"
19188            );
19189        }
19190
19191        // The fixture is not degenerate: the sinks and the window each
19192        // change the answer.
19193        let last = |p: &mut Pipeline| {
19194            p.clear_sequence_state();
19195            let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19196            p.logits_from_hidden(&hb[11 * hs..12 * hs])
19197        };
19198        let base = last(&mut p);
19199        let mut no_sinks = mimo_test_pipeline();
19200        for l in &mut no_sinks.kv_cache.layers {
19201            l.sinks = None;
19202        }
19203        assert_ne!(
19204            f32_bits(&last(&mut no_sinks)),
19205            f32_bits(&base),
19206            "sinks are live"
19207        );
19208        let mut wide = mimo_test_pipeline();
19209        wide.swa = Some((64, usize::MAX));
19210        assert_ne!(
19211            f32_bits(&last(&mut wide)),
19212            f32_bits(&base),
19213            "window is live"
19214        );
19215
19216        // Generation runs end to end on the same stack.
19217        p.clear_sequence_state();
19218        p.ignore_eos = true;
19219        let r = p.generate_from_ids(&ids, 4, None, None).unwrap();
19220        assert_eq!(r.token_ids.len(), 4);
19221    }
19222
19223    /// A synthetic MiMo draft stack of `n` layers for `mimo_test_pipeline`
19224    /// (the SWA geometry of its sliding layers: 2 KV heads, head 8 / V 4).
19225    fn mimo_test_mtp(n: usize, gain: f32) -> mimo_mtp::MimoMtp {
19226        let (hs, inter, nh, hd, vd, nkv) = (16usize, 24usize, 4usize, 8usize, 4usize, 2usize);
19227        let synth = |len: usize, salt: usize| -> Vec<f32> {
19228            (0..len)
19229                .map(|i| (((i * 37 + salt * 13 + 3) % 89) as f32 / 89.0 - 0.5) * 0.6 * gain)
19230                .collect()
19231        };
19232        let qt = |rows: usize, cols: usize, salt: usize| {
19233            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19234        };
19235        let layers = (0..n)
19236            .map(|k| {
19237                let s = 500 + k * 40;
19238                let mut kv = crate::kv_cache::LayerKvCache::new(nkv, hd);
19239                kv.sinks = Some(vec![0.3, -0.7, 1.1, 0.0]);
19240                MtpModule {
19241                    enorm: vec![1.0; hs],
19242                    hnorm: vec![1.0; hs],
19243                    eh_proj: qt(hs, 2 * hs, s),
19244                    layer: LayerWeights {
19245                        input_norm: vec![1.0; hs],
19246                        post_norm: vec![1.0; hs],
19247                        attn_out_norm: None,
19248                        ffn_out_norm: None,
19249                        layer_scale: None,
19250                        attn: AttnKind::Full {
19251                            wq: qt(nh * hd, hs, s + 1),
19252                            wk: qt(nkv * hd, hs, s + 2),
19253                            wv: qt(nkv * vd, hs, s + 3),
19254                            wo: qt(hs, nh * vd, s + 4),
19255                            q_norm: None,
19256                            k_norm: None,
19257                            output_gate: false,
19258                            softplus_gate: None,
19259                            bias: None,
19260                        },
19261                        ffn: FfnKind::Dense(DenseFfn {
19262                            gate_proj: qt(inter, hs, s + 5),
19263                            up_proj: qt(inter, hs, s + 6),
19264                            down_proj: qt(hs, inter, s + 7),
19265                            act: Act::Silu,
19266                            down_t: None,
19267                            segs: Vec::new(),
19268                        }),
19269                    },
19270                    final_norm: vec![1.0; hs],
19271                    kv,
19272                }
19273            })
19274            .collect();
19275        mimo_mtp::MimoMtp::from_layers(layers)
19276    }
19277
19278    fn mimo_greedy(p: &mut Pipeline, ids: &[u32], n: usize, spec: bool) -> GenerateResult {
19279        p.clear_sequence_state();
19280        p.speculative = spec;
19281        p.ignore_eos = true;
19282        p.sampler_config.temperature = 0.0;
19283        p.generate_from_ids(ids, n, None, None).unwrap()
19284    }
19285
19286    /// The draft stack's incremental rounds (a few rows per layer, last
19287    /// round's provisional rows dropped) give exactly the teacher-forced
19288    /// table of one causal pass per layer over the whole sequence — the
19289    /// table `tools/mimo_ref.py mtp` computes for variant A: layer k, row
19290    /// j reads (x[j+k+1], norm(h_j)) at RoPE position j.
19291    #[test]
19292    fn mimo_mtp_incremental_rounds_equal_the_teacher_forced_table() {
19293        // Both readings of the backbone hidden: pre-final-norm (default)
19294        // and post-final-norm (`CMF_MIMO_MTP_HIDDEN=post`).
19295        for post in [false, true] {
19296            let mut p = mimo_test_pipeline();
19297            // A non-trivial final norm, so the two readings differ.
19298            p.weights.final_norm = (0..p.hidden_size).map(|i| 0.5 + 0.1 * i as f32).collect();
19299            let mut st0 = mimo_test_mtp(3, 1.0);
19300            st0.post_norm_hidden = post;
19301            p.mimo_mtp = Some(st0);
19302            let ids: Vec<u32> = (0..14u32).map(|i| (i * 11 + 5) % 64).collect();
19303            let hs = p.hidden_size;
19304            let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19305            p.mimo_note_rows(&hb, 0);
19306            let mut st = p.mimo_mtp.take().unwrap();
19307            // Incremental: one round per t through the decode path (later
19308            // tokens from `ids`, the probe's teacher forcing).
19309            let k = 3;
19310            let mut inc = Vec::new();
19311            for t in 0..ids.len() - k - 1 {
19312                inc.push(p.mimo_mtp_draft(&mut st, t, &ids, k));
19313            }
19314            // Reference: per layer, ONE batched causal pass over all rows
19315            // with fresh caches.
19316            let s = ids.len();
19317            let mut reference = vec![vec![0u32; k]; s - k - 1];
19318            let mut fresh = mimo_test_mtp(3, 1.0);
19319            for (layer, m) in fresh.layers.iter_mut().enumerate() {
19320                let n = s - layer - 1;
19321                let mut cats = vec![0.0f32; n * 2 * hs];
19322                for j in 0..n {
19323                    let e = p.embed_single(ids[j + layer + 1]);
19324                    let raw = &hb[j * hs..(j + 1) * hs];
19325                    let g = if post {
19326                        inference::rms_norm(raw, &p.weights.final_norm, p.rms_eps, p.norm_style)
19327                    } else {
19328                        raw.to_vec()
19329                    };
19330                    let (ce, ch) = cats[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
19331                    inference::rms_norm_into(&e, &m.enorm, p.rms_eps, p.norm_style, ce);
19332                    inference::rms_norm_into(&g, &m.hnorm, p.rms_eps, p.norm_style, ch);
19333                }
19334                let mut x = vec![0.0f32; n * hs];
19335                m.eh_proj.matmat(&cats, n, &mut x, None);
19336                p.mimo_mtp_block(m, &mut x, n, 0);
19337                for (t, row) in reference.iter_mut().enumerate() {
19338                    let y = inference::rms_norm(
19339                        &x[t * hs..(t + 1) * hs],
19340                        &m.final_norm,
19341                        p.rms_eps,
19342                        p.norm_style,
19343                    );
19344                    row[layer] = sampler::argmax(&p.lm_head_forward(&y));
19345                }
19346            }
19347            assert_eq!(inc, reference, "post_norm_hidden = {post}");
19348            // Not a degenerate table: the drafts vary.
19349            let distinct: std::collections::HashSet<u32> =
19350                inc.iter().flatten().copied().collect();
19351            assert!(distinct.len() > 3, "{inc:?}");
19352            // Each layer's cache ends holding rows up to the last round start.
19353            let last_t = ids.len() - k - 2;
19354            for m in &st.layers {
19355                assert_eq!(m.kv.seq_len, last_t + 1);
19356            }
19357        }
19358    }
19359
19360    /// Greedy with the MiMo draft stack is the plain greedy stream, token
19361    /// for token — with the real draft layers (low acceptance) and with a
19362    /// drafter that is right most of the time (exercises accepted prefixes
19363    /// of every length, the KV truncation of the rejected rows and the
19364    /// logits hand-off to the loop top), under the default repetition
19365    /// penalty.
19366    #[test]
19367    fn mimo_speculative_greedy_equals_plain_greedy() {
19368        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19369        let ids: Vec<u32> = (0..9u32).map(|i| (i * 7 + 3) % 64).collect();
19370        let n = 24;
19371        let mut p = mimo_test_pipeline();
19372        let plain = mimo_greedy(&mut p, &ids, n, false);
19373        assert_eq!(plain.mtp_drafted, 0);
19374        assert_eq!(plain.token_ids.len(), n);
19375        let plain_kv = p.kv_cache.layers[0].seq_len;
19376
19377        // Real draft layers.
19378        p.mimo_mtp = Some(mimo_test_mtp(3, 1.0));
19379        let spec = mimo_greedy(&mut p, &ids, n, true);
19380        assert!(spec.mtp_drafted > 0, "the round must draft");
19381        assert_eq!(spec.token_ids, plain.token_ids);
19382        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19383
19384        // A drafter reading the true continuation with every fifth token
19385        // wrong: accepted prefixes of 0..=3 all occur.
19386        let mut truth: Vec<u32> = ids.clone();
19387        truth.extend(&plain.token_ids);
19388        let mut noisy = truth.clone();
19389        for (i, t) in noisy.iter_mut().enumerate() {
19390            if i % 5 == 0 {
19391                *t = (*t + 1) % 64;
19392            }
19393        }
19394        let mut st = mimo_test_mtp(3, 1.0);
19395        st.draft_override = Some(noisy);
19396        p.mimo_mtp = Some(st);
19397        let spec = mimo_greedy(&mut p, &ids, n, true);
19398        assert_eq!(spec.token_ids, plain.token_ids);
19399        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19400        let stats = p.mimo_mtp.as_ref().unwrap().stats.clone();
19401        assert_eq!(stats.accepted as usize, spec.mtp_accepted);
19402        assert!(spec.mtp_accepted > 0 && spec.mtp_accepted < spec.mtp_drafted);
19403        assert!(
19404            stats.accept_hist.iter().filter(|&&c| c > 0).count() >= 3,
19405            "{:?}",
19406            stats.accept_hist
19407        );
19408        assert!(stats.tokens_per_round() > 1.5, "{}", stats.line());
19409
19410        // A perfect drafter: every draft accepted, rounds of K+1 tokens,
19411        // and the budget is never overrun.
19412        let mut st = mimo_test_mtp(3, 1.0);
19413        st.draft_override = Some(truth);
19414        p.mimo_mtp = Some(st);
19415        let spec = mimo_greedy(&mut p, &ids, n, true);
19416        assert_eq!(spec.token_ids, plain.token_ids);
19417        assert_eq!(spec.mtp_accepted, spec.mtp_drafted);
19418        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19419
19420        // CMF_MTP=0 path: the stack is attached but idle.
19421        let off = mimo_greedy(&mut p, &ids, n, false);
19422        assert_eq!(off.token_ids, plain.token_ids);
19423        assert_eq!(off.mtp_drafted, 0);
19424    }
19425
19426    /// The wgpu graphs carry MiMo-V2's attention per layer (KV heads,
19427    /// narrow V, sinks, windows, two RoPE tables): no attention-level
19428    /// decline for it any more, and the geometry each layer hands the
19429    /// graph is exactly what the CPU attention reads for that layer. The
19430    /// descriptive reasons stay (the Metal graphs and the q1 dropin still
19431    /// decline on them), and what the per-layer geometry cannot express
19432    /// keeps a named wgpu decline.
19433
19434    #[test]
19435    fn mimo_shaped_model_rides_the_wgpu_graph_geometry() {
19436        let p = mimo_test_pipeline();
19437        assert_eq!(
19438            p.graph_attn_decline_reason(),
19439            Some("per-layer KV head counts")
19440        );
19441        assert_eq!(p.wgpu_graph_attn_decline(), None);
19442        let g0 = p.graph_attn_geom(0).expect("full layer geometry");
19443        assert_eq!(
19444            (g0.nkv, g0.dv, g0.rd, g0.window, g0.sink.is_some()),
19445            (1, 4, 4, None, false)
19446        );
19447        assert_eq!(g0.invf, p.inv_freq.as_slice());
19448        let g1 = p.graph_attn_geom(1).expect("sliding layer geometry");
19449        assert_eq!((g1.nkv, g1.dv, g1.rd, g1.window), (2, 4, 4, Some(3)));
19450        assert_eq!(g1.sink, Some(&[0.5f32, -1.0, 1.5, 0.0][..]));
19451        assert_eq!(g1.invf, p.inv_freq_local.as_ref().unwrap().as_slice());
19452        assert_ne!(g0.invf, g1.invf, "two RoPE tables");
19453        let g3 = p.graph_attn_geom(3).expect("full layer geometry");
19454        assert_eq!((g3.nkv, g3.window, g3.sink.is_some()), (1, None, false));
19455
19456        // No wgpu device in this process: the builders run and decline on
19457        // the (f32, unmapped) experts — never with an attention line.
19458        let emb = p.embed_single(3);
19459        let mut lg = Vec::new();
19460        assert!(
19461            p.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, p.num_layers)
19462                .is_none()
19463        );
19464        let mut hid = emb.clone();
19465        assert_eq!(
19466            p.try_batch_graph_wgpu(&mut hid, &[0], 1, None),
19467            crate::gpu::BatchGraphOutcome::Declined
19468        );
19469        assert_eq!(hid, emb, "a declined batch graph leaves the rows untouched");
19470        assert!(p.try_multi_burst(3, 0, 4).is_none());
19471        assert!(
19472            p.graph_declines().is_empty(),
19473            "no attention decline logged: {:?}",
19474            p.graph_declines()
19475        );
19476        // (No assertion on graph_prefill_preferred: with no attention
19477        // decline it follows the device — a test process that brought a
19478        // wgpu adapter up routes this resident MoE through the graph.)
19479
19480        let plain = || create_test_pipeline(8, 16, 2, 1, 4, 2, 32);
19481        assert_eq!(plain().graph_attn_decline_reason(), None);
19482        assert_eq!(plain().wgpu_graph_attn_decline(), None);
19483        assert!(
19484            plain().graph_attn_geom(0).is_none(),
19485            "uniform models keep the historical arms"
19486        );
19487        let mut q = plain();
19488        q.set_layer_sinks(1, vec![0.25, -0.25]).unwrap();
19489        assert_eq!(
19490            q.graph_attn_decline_reason(),
19491            Some("learned attention sinks")
19492        );
19493        assert_eq!(
19494            q.graph_attn_geom(1).unwrap().sink,
19495            Some(&[0.25f32, -0.25][..])
19496        );
19497        let mut q = plain();
19498        q.set_attn_geometry(None, Some(2)).unwrap();
19499        assert_eq!(
19500            q.graph_attn_decline_reason(),
19501            Some("V heads narrower than Q/K heads")
19502        );
19503        assert_eq!(q.graph_attn_geom(0).unwrap().dv, 2);
19504        let mut q = plain();
19505        q.sliding_layers = Some(vec![true, false]);
19506        q.swa = Some((4, usize::MAX));
19507        assert_eq!(q.graph_attn_decline_reason(), Some("sliding-window layers"));
19508        assert_eq!(q.graph_attn_geom(0).unwrap().window, Some(4));
19509        assert_eq!(q.graph_attn_geom(1).unwrap().window, None);
19510
19511        // Outside the per-layer geometry: a named wgpu decline, logged
19512        // once per site.
19513        let mut q = mimo_test_pipeline();
19514        q.rope_scale = 2.0;
19515        assert_eq!(
19516            q.wgpu_graph_attn_decline(),
19517            Some("scaled RoPE positions with per-layer geometry")
19518        );
19519        let emb = q.embed_single(3);
19520        assert!(
19521            q.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, q.num_layers)
19522                .is_none()
19523        );
19524        let _ = q.try_token_graph_wgpu_steps(&emb, 1, &mut lg, 1, None, None, 0, q.num_layers);
19525        let lines = q.graph_declines();
19526        assert_eq!(
19527            lines
19528                .iter()
19529                .filter(|l| l.starts_with("wgpu token graph") && l.contains("scaled RoPE"))
19530                .count(),
19531            1,
19532            "{lines:?}"
19533        );
19534    }
19535
19536    #[test]
19537    fn mimo_verify_rewind_preserves_lagging_host_caches() {
19538        let mut p = mimo_test_pipeline();
19539        for (li, layer) in p.kv_cache.layers.iter_mut().enumerate() {
19540            let row = vec![0.0; layer.num_kv_heads * layer.head_dim];
19541            for _ in 0..if li == 0 { 2 } else { 12 } {
19542                layer.append(&row, &row, &[]);
19543            }
19544        }
19545        p.mimo_verify_rewind(9).unwrap();
19546        assert_eq!(p.kv_cache.layers[0].seq_len, 2);
19547        for layer in &p.kv_cache.layers[1..] {
19548            assert_eq!(layer.seq_len, 9);
19549        }
19550    }
19551
19552    /// CMF_LAYER_DUMP: the decode walk and the batched prefill both write
19553    /// every (position, layer) hidden, the two sets agree byte for byte,
19554    /// and the last layer's file is the stack output.
19555    #[test]
19556    fn layer_dump_covers_every_position_and_layer_on_both_walks() {
19557        let dir = std::env::temp_dir().join(format!("cmf-layer-dump-{}", std::process::id()));
19558        let _ = std::fs::remove_dir_all(&dir);
19559        let mut p = mimo_test_pipeline();
19560        let hs = p.hidden_size;
19561        let ids = [5u32, 9, 11, 2, 40];
19562        p.layer_dump = Some(dir.join("decode"));
19563        for (pos, &id) in ids.iter().enumerate() {
19564            let e = p.embed_single(id);
19565            let _ = p.forward_layers(&e, pos, None);
19566        }
19567        p.clear_sequence_state();
19568        p.layer_dump = Some(dir.join("prefill"));
19569        let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19570        for pos in 0..ids.len() {
19571            for li in 0..p.num_layers {
19572                let name = format!("p{pos:06}_l{li:02}.f32");
19573                let a = std::fs::read(dir.join("decode").join(&name)).unwrap();
19574                let b = std::fs::read(dir.join("prefill").join(&name)).unwrap();
19575                assert_eq!(a.len(), hs * 4, "{name}");
19576                assert_eq!(a, b, "{name}");
19577            }
19578        }
19579        let last = std::fs::read(dir.join("prefill").join("p000004_l03.f32")).unwrap();
19580        let vals: Vec<f32> = last
19581            .chunks(4)
19582            .map(|c| f32::from_le_bytes(c.try_into().unwrap()))
19583            .collect();
19584        assert_eq!(f32_bits(&vals), f32_bits(&hb[4 * hs..5 * hs]));
19585        let _ = std::fs::remove_dir_all(&dir);
19586    }
19587
19588    #[test]
19589    fn attn_geometry_and_sinks_are_validated() {
19590        let mut p = create_test_pipeline(8, 16, 4, 2, 4, 2, 32);
19591        assert!(
19592            p.set_attn_geometry(Some(vec![2]), None).is_err(),
19593            "one entry per layer"
19594        );
19595        assert!(
19596            p.set_attn_geometry(Some(vec![2, 3]), None).is_err(),
19597            "3 does not divide 4"
19598        );
19599        assert!(p.set_attn_geometry(Some(vec![2, 0]), None).is_err());
19600        assert!(p.set_attn_geometry(None, Some(0)).is_err());
19601        assert!(
19602            p.set_attn_geometry(None, Some(5)).is_err(),
19603            "V wider than the head"
19604        );
19605        p.set_attn_geometry(None, Some(4)).unwrap();
19606        assert_eq!(
19607            p.v_head_dim, None,
19608            "v_head_dim == head_dim is the uniform case"
19609        );
19610        p.set_layer_sinks(1, vec![0.1; 4]).unwrap();
19611        p.set_attn_geometry(Some(vec![1, 4]), None).unwrap();
19612        assert_eq!(p.kv_cache.layers[0].num_kv_heads, 1);
19613        assert_eq!(p.kv_cache.layers[1].num_kv_heads, 4);
19614        assert!(
19615            p.kv_cache.layers[1].sinks.is_some(),
19616            "a reshape keeps the layer's sinks"
19617        );
19618        assert_eq!(p.layer_geom(1).0, 4);
19619        assert!(
19620            p.set_layer_sinks(0, vec![0.0; 3]).is_err(),
19621            "one sink per Q head"
19622        );
19623        assert!(p.set_layer_sinks(7, vec![0.0; 4]).is_err());
19624        assert!(p.set_layer_sinks(0, vec![f32::NAN, 0.0, 0.0, 0.0]).is_err());
19625    }
19626
19627    /// The O(1) Nyström state replaces a plain full-context softmax; it
19628    /// must never be armed on a sliding, sink or narrow-V layer.
19629    #[test]
19630    fn o1_is_never_armed_on_sink_window_or_narrow_v_layers() {
19631        let cfg = || {
19632            Some(crate::nystrom::O1Cfg {
19633                layers: crate::nystrom::O1Layers::All,
19634                m: 4,
19635                w: 8,
19636                sink: 2,
19637                rect: crate::nystrom::O1Rect::Aggregate,
19638            })
19639        };
19640        let mut p = mimo_test_pipeline();
19641        p.set_o1(cfg());
19642        assert!(!p.o1_active(), "every MiMo-shaped layer is ineligible");
19643        let mut q = create_test_pipeline(8, 16, 2, 1, 4, 3, 64);
19644        q.set_layer_sinks(1, vec![0.0, 0.0]).unwrap();
19645        q.sliding_layers = Some(vec![false, false, true]);
19646        q.swa = Some((4, usize::MAX));
19647        q.set_o1(cfg());
19648        assert_eq!(q.o1_flags, vec![true, false, false]);
19649    }
19650
19651    #[test]
19652    fn o1_batch_transition_publishes_one_epoch_before_serial_handoff() {
19653        const B: usize = 19;
19654        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19655        p.set_o1(Some(crate::nystrom::O1Cfg {
19656            layers: crate::nystrom::O1Layers::All,
19657            m: 4,
19658            w: 8,
19659            sink: 2,
19660            rect: crate::nystrom::O1Rect::Aggregate,
19661        }));
19662        p.o1_begin_with_prefix(Some(B));
19663        let ids: Vec<u32> = (0..B as u32).collect();
19664        let _ = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19665
19666        assert_eq!(p.o1_epoch, 1, "all layers publish one completed transition");
19667        assert!(p.kv_cache.layers.iter().all(|l| l.o1_sealed()));
19668        let next = p.embed_single(B as u32);
19669        let _ = p.forward_layers(&next, B, None);
19670        assert_eq!(p.o1_epoch, 1, "sealed handoff must not republish the epoch");
19671    }
19672
19673    #[test]
19674    fn o1_pair_transition_commits_scratch_before_epoch_publication() {
19675        const B: usize = 19;
19676        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19677        // Keep a real recurrent layer ahead of the Full O(1) layer so the
19678        // pair test observes the GDN lane-2 scratch swap at the same
19679        // boundary, rather than only exercising an artificial scratch vec.
19680        let gdn_cfg = crate::linear_core::GdnCfg {
19681            num_v_heads: 2,
19682            num_k_heads: 1,
19683            key_head_dim: 2,
19684            value_head_dim: 4,
19685            conv_kernel: 3,
19686            hidden_size: 8,
19687            rms_eps: 1e-6,
19688            output_gate_sigmoid: false,
19689        };
19690        let synth = |n: usize, salt: usize| -> Vec<f32> {
19691            (0..n)
19692                .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
19693                .collect()
19694        };
19695        let qt = |rows: usize, cols: usize, salt: usize| {
19696            crate::qtensor::QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19697        };
19698        let c_dim = gdn_cfg.conv_dim();
19699        let vd = gdn_cfg.num_v_heads * gdn_cfg.value_head_dim;
19700        p.weights.layers[0].attn = AttnKind::LinearGdn(crate::linear_core::GdnWeights {
19701            in_proj_qkv: qt(c_dim, 8, 1),
19702            in_proj_z: qt(vd, 8, 2),
19703            in_proj_a: qt(gdn_cfg.num_v_heads, 8, 3),
19704            in_proj_b: qt(gdn_cfg.num_v_heads, 8, 4),
19705            conv1d: synth(c_dim * gdn_cfg.conv_kernel, 5),
19706            a_log: vec![0.2, 0.5],
19707            dt_bias: synth(gdn_cfg.num_v_heads, 6),
19708            norm: vec![1.0; gdn_cfg.value_head_dim],
19709            out_proj: qt(8, vd, 7),
19710        });
19711        p.gdn_cfg = Some(gdn_cfg);
19712        p.set_o1(Some(crate::nystrom::O1Cfg {
19713            layers: crate::nystrom::O1Layers::All,
19714            m: 4,
19715            w: 8,
19716            sink: 2,
19717            rect: crate::nystrom::O1Rect::Aggregate,
19718        }));
19719        p.o1_begin_with_prefix(Some(B));
19720        for pos in 0..B - 2 {
19721            let emb = p.embed_single(pos as u32);
19722            let _ = p.forward_layers(&emb, pos, None);
19723        }
19724        let lane1_state = p.kv_cache.layers[0].linear_state.clone();
19725
19726        let e1 = p.embed_single((B - 2) as u32);
19727        let e2 = p.embed_single((B - 1) as u32);
19728        let _ = p.forward_pair(&e1, &e2, B - 2);
19729
19730        assert_eq!(p.o1_epoch, 1, "pair crossing B publishes one epoch");
19731        assert!(
19732            p.kv_cache
19733                .layers
19734                .iter()
19735                .enumerate()
19736                .all(|(li, l)| !p.o1_flags[li] || l.o1_sealed())
19737        );
19738        assert!(!p.kv_cache.layers[0].linear_state.is_empty());
19739        assert_ne!(
19740            p.kv_cache.layers[0].linear_state, lane1_state,
19741            "real pair must commit GDN lane 2 before returning"
19742        );
19743        assert!(p.kv_cache.layers[0].linear_scratch.is_empty());
19744        let next = p.embed_single(B as u32);
19745        let _ = p.forward_layers(&next, B, None);
19746        assert_eq!(p.o1_epoch, 1, "serial continuation must reuse the epoch");
19747    }
19748
19749    #[test]
19750    fn o1_error_observation_stays_terminal_until_reset() {
19751        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19752        p.set_o1(Some(crate::nystrom::O1Cfg {
19753            layers: crate::nystrom::O1Layers::All,
19754            m: 4,
19755            w: 8,
19756            sink: 2,
19757            rect: crate::nystrom::O1Rect::Aggregate,
19758        }));
19759        p.o1_begin();
19760        p.kv_cache.layers[0].o1_abort("synthetic transition failure".into());
19761
19762        assert!(p.o1_seal_checked().is_err());
19763        assert!(
19764            p.o1_seal_checked().is_err(),
19765            "retry must see the sticky error"
19766        );
19767        let k = vec![0.2f32; 4];
19768        let v = vec![0.3f32; 4];
19769        p.kv_cache.layers[0].append(&k, &v, &[]);
19770        assert_eq!(p.kv_cache.layers[0].seq_len, 0);
19771
19772        p.reset_session();
19773        p.o1_begin();
19774        p.kv_cache.layers[0].append(&k, &v, &[]);
19775        assert_eq!(p.kv_cache.layers[0].seq_len, 1);
19776    }
19777
19778    #[test]
19779    fn nll_graph_failure_is_terminal_and_request_is_reusable() {
19780        let ids = vec![1u32, 2, 3, 4, 5, 6];
19781        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19782        p.graph_logits = Some(vec![123.0]);
19783        p.graph_want_logits = true;
19784        p.graph_failed
19785            .store(true, std::sync::atomic::Ordering::Relaxed);
19786        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
19787        let err = p.nll_ids_from(&ids, 0).expect_err("prior graph failure");
19788        assert!(err.contains("before NLL"));
19789        assert!(p.graph_logits.is_none());
19790        assert!(!p.graph_want_logits);
19791        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19792        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19793
19794        let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19795        let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
19796        let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
19797        assert_eq!(actual.1, expected.1);
19798        assert!((actual.0 - expected.0).abs() < 1e-9);
19799    }
19800
19801    #[test]
19802    fn nll_forward_failure_discards_partial_score_and_clears_sidechannels() {
19803        let ids = vec![1u32, 2, 3, 4, 5, 6];
19804        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19805        p.nll_test_fail_at = Some(1);
19806        let err = p
19807            .nll_ids_from(&ids, 0)
19808            .expect_err("one-shot forward failure");
19809        assert!(err.contains("forward") || err.contains("score row"));
19810        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19811        assert!(!p.graph_want_logits);
19812        assert!(p.graph_logits.is_none());
19813        assert!(p.kv_history.is_empty());
19814
19815        let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19816        let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
19817        let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
19818        assert_eq!(actual.1, expected.1);
19819        assert!((actual.0 - expected.0).abs() < 1e-9);
19820    }
19821
19822    #[test]
19823    fn nll_serial_failure_before_first_row_is_reported() {
19824        let ids = vec![1u32, 2, 3, 4];
19825        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19826        p.nll_test_force_serial = true;
19827        p.nll_test_fail_at = Some(0);
19828        let err = p.nll_ids_from(&ids, 0).expect_err("serial forward failure");
19829        assert!(err.contains("serial forward"));
19830        assert!(p.kv_history.is_empty());
19831        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19832        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19833    }
19834
19835    #[test]
19836    fn ffn_probe_failure_discards_recorder_and_state() {
19837        let ids = vec![1u32, 2, 3, 4];
19838        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19839        p.nll_test_fail_at = Some(0);
19840        let err = p
19841            .probe_ffn_mass_batch(&ids)
19842            .expect_err("probe forward failure");
19843        assert!(err.contains("NLL"));
19844        assert!(FFN_PROBE.with(|probe| probe.borrow().is_none()));
19845        assert!(p.kv_history.is_empty());
19846        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19847    }
19848
19849    #[test]
19850    fn nll_test_controls_are_pipeline_scoped() {
19851        let ids = vec![1u32, 2, 3, 4];
19852        let mut failing = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19853        let mut unaffected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19854        failing.nll_test_force_serial = true;
19855        failing.nll_test_fail_at = Some(0);
19856
19857        assert!(!failing.can_prefill_batched());
19858        assert!(unaffected.can_prefill_batched());
19859        let expected = unaffected
19860            .nll_ids_from(&ids, 0)
19861            .expect("unaffected pipeline remains usable");
19862        let err = failing
19863            .nll_ids_from(&ids, 0)
19864            .expect_err("failure injection belongs to failing pipeline");
19865        assert!(err.contains("serial forward"));
19866        assert!(failing.nll_test_fail_at.is_none());
19867        assert!(unaffected.can_prefill_batched());
19868        let actual = unaffected
19869            .nll_ids_from(&ids, 0)
19870            .expect("unaffected pipeline remains reusable");
19871        assert_eq!(actual.1, expected.1);
19872        assert!((actual.0 - expected.0).abs() < 1e-9);
19873    }
19874
19875    #[test]
19876    fn forward_ids_failure_channel_is_terminal_and_reusable() {
19877        let ids = vec![1u32, 2, 3, 4, 5, 6];
19878        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19879        p.graph_logits = Some(vec![123.0]);
19880        p.graph_want_logits = true;
19881        p.graph_failed
19882            .store(true, std::sync::atomic::Ordering::Relaxed);
19883        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
19884
19885        let err = p
19886            .forward_ids(&ids, None)
19887            .expect_err("a failed forward must not become a valid head result");
19888        assert!(err.contains("forward_ids setup"));
19889        assert!(p.graph_logits.is_none());
19890        assert!(!p.graph_want_logits);
19891        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19892        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19893        assert_eq!(p.kv_cache.seq_len(), 0);
19894
19895        let expected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64)
19896            .forward_ids(&ids, None)
19897            .expect("fresh forward_ids");
19898        let actual = p
19899            .forward_ids(&ids, None)
19900            .expect("pipeline remains reusable after a failed forward");
19901        assert_eq!(actual.len(), expected.len());
19902        assert!(
19903            actual
19904                .iter()
19905                .zip(expected)
19906                .all(|(a, b)| (a - b).abs() < 1e-9)
19907        );
19908        assert_eq!(p.kv_cache.seq_len(), ids.len());
19909    }
19910
19911    #[test]
19912    fn sigmoid_router_floor_is_explicit_per_architecture() {
19913        // GLM-5's noaux_tc reference uses +1e-20 while the generic
19914        // LFM2-compatible path uses +1e-6.  At low (but representable)
19915        // sigmoid scores, silently sharing the latter changes expert weights
19916        // by orders of magnitude and can make a routed layer look coherent
19917        // while discarding its expert contribution.
19918        let zero = || DenseFfn {
19919            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19920            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19921            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19922            act: Act::Silu,
19923            down_t: None,
19924            segs: Vec::new(),
19925        };
19926        let m = MoeFfn {
19927            router: QTensor::from_f32(vec![0.0; 4], 2, 2),
19928            experts: vec![zero(), zero()],
19929            top_k: 1,
19930            norm_topk_prob: true,
19931            router_sigmoid: true,
19932            expert_bias: None,
19933            routed_scaling: 2.5,
19934            route_tau: None,
19935            shared: None,
19936            stats: std::cell::RefCell::new(Vec::new()),
19937            act_sq: std::cell::RefCell::new(Vec::new()),
19938            act_rows: std::cell::RefCell::new(Vec::new()),
19939            mask: None,
19940            per_expert_scale: None,
19941            router_input_norm: false,
19942            resonance: None,
19943            grown: Vec::new(),
19944        };
19945        let logits = [-20.0f32, -20.0];
19946        let (_, p, glm_wsum) = moe_route_with_eps(&logits, &m, None, 1e-20);
19947        let (_, _, generic_wsum) = moe_route(&logits, &m, None);
19948        let expected = (p[0] + 1e-20) / m.routed_scaling;
19949        assert!((glm_wsum - expected).abs() < 1e-15);
19950        assert!(generic_wsum > glm_wsum * 100.0);
19951    }
19952
19953    #[test]
19954    fn resonance_scores_match_formula_and_stable_tie() {
19955        let r = Resonance {
19956            // Three descriptors, hidden=2, one projection row each.
19957            mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0],
19958            u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
19959            k: 1,
19960            bias: vec![1.5, 0.5, 0.0],
19961            shell: Vec::new(),
19962        };
19963        let x = [1.0f32, 1.0];
19964        let mut got = vec![0.0; 3];
19965        r.scores(&x, &mut got);
19966        // Expert 0 and 1 are an exact score tie; the CPU top-1 contract uses
19967        // the lower index.  The values also check d² - (U·d)², not just tie
19968        // ordering.
19969        assert!((got[0] - 0.5).abs() < 1e-6);
19970        assert!((got[1] - 0.5).abs() < 1e-6);
19971        assert!(got[2].abs() < 1e-6);
19972        let best = got
19973            .iter()
19974            .enumerate()
19975            .max_by(|(ia, a), (ib, b)| a.partial_cmp(b).unwrap().then(ib.cmp(ia)))
19976            .map(|(i, _)| i);
19977        assert_eq!(best, Some(0));
19978        assert!(got.iter().all(|v| v.is_finite()));
19979    }
19980
19981    /// The growth shell (spec §2): a grown expert whose reconstruction
19982    /// error lies outside its shell scores −∞, one inside keeps the exact
19983    /// resonance score, trunk rows (`+inf` shell) are bit-identical to the
19984    /// shell-less computation; the process-wide switch disables it.
19985    #[test]
19986    fn resonance_shell_masks_outside_keeps_inside_and_trunk_bits() {
19987        // hidden = 2, rank 1. Experts 0/1 = trunk (shell +inf); 2 and 3 =
19988        // grown, the same descriptor (μ = (0, 1), u = (1, 1)) with shells
19989        // 6.0 and 0.25. At x' = (3, 0): d = (3, −1), d² = 10, proj =
19990        // (3 − 1)² = 4, err = 6 exactly — on the boundary of expert 2's
19991        // shell (kept: the rule is strict `>`), outside expert 3's.
19992        let plain = Resonance {
19993            mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 1.0],
19994            u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0, 1.0],
19995            k: 1,
19996            bias: vec![1.5, 0.5, 0.0, 0.0],
19997            shell: Vec::new(),
19998        };
19999        let shelled = Resonance {
20000            mu: plain.mu.clone(),
20001            u: plain.u.clone(),
20002            k: 1,
20003            bias: plain.bias.clone(),
20004            shell: vec![f32::INFINITY, f32::INFINITY, 6.0, 0.25],
20005        };
20006        assert!(!plain.has_shell());
20007        assert!(shelled.has_shell());
20008        let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
20009        let (mut a, mut b) = (vec![0.0; 4], vec![0.0; 4]);
20010        set_growth_shell(Some(true));
20011        assert!(growth_shell_enabled());
20012        // x = (1, 1): the grown experts reconstruct it exactly (err 0):
20013        // inside both shells, every row the shell-less bits.
20014        let x = [1.0f32, 1.0];
20015        plain.scores(&x, &mut a);
20016        shelled.scores(&x, &mut b);
20017        assert_eq!(bits(&a), bits(&b), "inside every shell: unchanged");
20018        assert!(a[2] == 0.0 && a[3] == 0.0);
20019        // x' = (3, 0): expert 3 → −∞, expert 2 (err == shell) and the
20020        // trunk rows keep their exact bits.
20021        let xo = [3.0f32, 0.0];
20022        plain.scores(&xo, &mut a);
20023        shelled.scores(&xo, &mut b);
20024        assert_eq!(a[2], -6.0);
20025        assert_eq!(a[3], -6.0);
20026        assert_eq!(bits(&a[..3]), bits(&b[..3]), "trunk rows + the boundary row");
20027        assert_eq!(b[3], f32::NEG_INFINITY, "outside the shell: −∞");
20028        assert_eq!(shelled.effective_shell(4), shelled.shell);
20029        // The switch (`CMF_GROWTH_SHELL=off` / `growth-eval --shell off`):
20030        // all +inf, the shell-less bits everywhere.
20031        set_growth_shell(Some(false));
20032        assert!(!growth_shell_enabled());
20033        assert_eq!(shelled.effective_shell(4), vec![f32::INFINITY; 4]);
20034        shelled.scores(&xo, &mut b);
20035        assert_eq!(bits(&a), bits(&b));
20036        set_growth_shell(None);
20037        // A shell vector shorter than the expert count masks nothing
20038        // beyond it (a legacy layer whose tail has no shell).
20039        let short = Resonance {
20040            shell: vec![f32::INFINITY, f32::INFINITY],
20041            ..shelled
20042        };
20043        set_growth_shell(Some(true));
20044        short.scores(&xo, &mut b);
20045        assert_eq!(bits(&a), bits(&b));
20046        set_growth_shell(None);
20047    }
20048
20049    /// `moe_route` with −∞ logits (a grown expert outside its shell):
20050    /// top-1 is the best finite expert with weight exactly 1.0 on both
20051    /// the softmax and the sigmoid path; all −∞ degrades to uniform.
20052    #[test]
20053    fn moe_route_neg_inf_logits_pick_the_best_finite_expert_with_weight_one() {
20054        let zero = || DenseFfn {
20055            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20056            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20057            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20058            act: Act::Silu,
20059            down_t: None,
20060            segs: Vec::new(),
20061        };
20062        let moe = |sigmoid: bool| MoeFfn {
20063            router: QTensor::from_f32(vec![0.0; 8], 4, 2),
20064            experts: vec![zero(), zero(), zero(), zero()],
20065            top_k: 1,
20066            norm_topk_prob: true,
20067            router_sigmoid: sigmoid,
20068            expert_bias: None,
20069            routed_scaling: 1.0,
20070            route_tau: None,
20071            shared: None,
20072            stats: std::cell::RefCell::new(Vec::new()),
20073            act_sq: std::cell::RefCell::new(Vec::new()),
20074            act_rows: std::cell::RefCell::new(Vec::new()),
20075            mask: None,
20076            per_expert_scale: None,
20077            router_input_norm: false,
20078            resonance: None,
20079            grown: Vec::new(),
20080        };
20081        let logits = [-1.0f32, f32::NEG_INFINITY, -0.5, f32::NEG_INFINITY];
20082        for sigmoid in [false, true] {
20083            let m = moe(sigmoid);
20084            let (idx, p, wsum) = moe_route(&logits, &m, None);
20085            assert_eq!(idx, vec![2], "sigmoid {sigmoid}");
20086            assert_eq!(p[1], 0.0);
20087            assert_eq!(p[3], 0.0);
20088            assert!(p[2] > p[0] && p[0] > 0.0);
20089            assert!(p.iter().all(|v| v.is_finite()));
20090            let w = p[2] / wsum;
20091            if sigmoid {
20092                // The sigmoid renorm keeps its reference floor (+1e-6).
20093                assert!((w - 1.0).abs() < 1e-5, "sigmoid: weight {w}");
20094            } else {
20095                assert_eq!(w, 1.0, "softmax: the renormalized top-1 weight is exactly 1");
20096            }
20097            // Masked experts stay masked even when they are the only ones
20098            // "admitted" by an allow-list that covers everything.
20099            let (idx, _, _) = moe_route(&logits, &m, Some(&[true, true, true, true]));
20100            assert_eq!(idx, vec![2]);
20101        }
20102        // A finite expert always beats −∞ whatever the bias / order.
20103        let m = moe(false);
20104        let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, f32::NEG_INFINITY, f32::NEG_INFINITY, -9.0], &m, None);
20105        assert_eq!(idx, vec![3]);
20106        // Every expert at −∞ (cannot happen on a grown file — trunk rows
20107        // have no shell): uniform, finite, lowest index.
20108        let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 4], &m, None);
20109        assert_eq!(idx, vec![0]);
20110        assert!(p.iter().all(|&v| v == 0.25));
20111        assert!(wsum.is_finite() && wsum > 0.0);
20112    }
20113
20114    /// The resonance router (top-1) selects by the raw score as the
20115    /// trainer and the graph do — not by softmax probabilities, where two
20116    /// scores closer than 2^-25 collapse to the same `exp(l − max) = 1.0`
20117    /// and the LOWER index wins a token whose score is strictly smaller.
20118    #[test]
20119    fn resonance_route_picks_the_raw_argmax_not_the_softmax_tie() {
20120        let zero = || DenseFfn {
20121            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20122            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20123            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20124            act: Act::Silu,
20125            down_t: None,
20126            segs: Vec::new(),
20127        };
20128        let moe = |resonant: bool, norm_topk: bool| MoeFfn {
20129            router: QTensor::from_f32(vec![0.0; 6], 3, 2),
20130            experts: vec![zero(), zero(), zero()],
20131            top_k: 1,
20132            norm_topk_prob: norm_topk,
20133            router_sigmoid: false,
20134            expert_bias: None,
20135            routed_scaling: 1.0,
20136            route_tau: None,
20137            shared: None,
20138            stats: std::cell::RefCell::new(Vec::new()),
20139            act_sq: std::cell::RefCell::new(Vec::new()),
20140            act_rows: std::cell::RefCell::new(Vec::new()),
20141            mask: None,
20142            per_expert_scale: None,
20143            router_input_norm: false,
20144            resonance: resonant.then(|| Resonance {
20145                mu: vec![0.0; 6],
20146                u: Vec::new(),
20147                k: 0,
20148                bias: vec![0.0; 3],
20149                shell: Vec::new(),
20150            }),
20151            grown: Vec::new(),
20152        };
20153        // lo = −0.1, hi = the next f32 towards zero: hi − lo = 2^-27 <
20154        // 2^-25, so exp(lo − hi) rounds to exactly 1.0 — a softmax tie.
20155        let lo = -0.1f32;
20156        let hi = f32::from_bits(lo.to_bits() - 1);
20157        assert!(hi > lo && hi - lo < 2f32.powi(-25));
20158        assert_eq!((lo - hi).exp(), 1.0, "the tie this test is about");
20159        // The gated MoE (softmax) path: the tie hands the token to index 0.
20160        let (idx, _, _) = moe_route(&[lo, hi], &moe(false, true), None);
20161        assert_eq!(idx, vec![0], "softmax tie → lower index (the gated contract)");
20162        // The resonance path: the strictly larger raw score wins, weight
20163        // exactly 1.0 with and without norm_topk.
20164        for norm in [true, false] {
20165            let m = moe(true, norm);
20166            let (idx, p, wsum) = moe_route(&[lo, hi, f32::NEG_INFINITY], &m, None);
20167            assert_eq!(idx, vec![1], "norm_topk {norm}");
20168            assert_eq!(p, vec![0.0, 1.0, 0.0]);
20169            assert_eq!(p[1] / wsum, 1.0);
20170            // An exact tie: the first maximum (as `resonance_winner` and
20171            // `embryo_core_route_pick`).
20172            let (idx, _, _) = moe_route(&[hi, hi, lo], &m, None);
20173            assert_eq!(idx, vec![0]);
20174            // `−∞` never wins; the admitted set is honoured.
20175            let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, lo, hi], &m, None);
20176            assert_eq!(idx, vec![2]);
20177            let (idx, p, wsum) = moe_route(&[lo, hi, hi], &m, Some(&[true, false, false]));
20178            assert_eq!(idx, vec![0]);
20179            assert_eq!(p[0] / wsum, 1.0);
20180            // Every admitted expert at −∞: the generic path's uniform
20181            // fallback (lowest index, finite weights).
20182            let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 3], &m, None);
20183            assert_eq!(idx, vec![0]);
20184            assert!(p.iter().all(|v| v.is_finite()) && wsum.is_finite() && wsum > 0.0);
20185        }
20186    }
20187}