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            // The device path takes several prompt tokens per layer frame;
4147            // the host path runs them one by one inside the same call.
4148            let end = (pos + crate::qwen4_exp::prefill_chunk()).min(input_ids.len());
4149            let want_logits = end == input_ids.len();
4150            let mut lg = Vec::new();
4151            if let Some(b) = &mut self.qwen4_exp {
4152                crate::qwen4_exp::forward_tokens(
4153                    &b.0,
4154                    &b.1,
4155                    &b.2,
4156                    &mut b.3,
4157                    &input_ids[pos..end],
4158                    pos,
4159                    &self.inv_freq,
4160                    self.pool.as_deref(),
4161                    &mut lg,
4162                    want_logits,
4163                );
4164            }
4165            if want_logits {
4166                self.graph_logits = Some(lg);
4167            }
4168            pos = end;
4169            hidden.fill(0.0);
4170        }
4171        while self.dsv4.is_some()
4172            && mtp.is_none()
4173            && pos < input_ids.len()
4174            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4175        {
4176            let end = (pos + prefill_chunk()).min(input_ids.len());
4177            let ids: Vec<u32> = input_ids[pos..end].to_vec();
4178            let mut lg = Vec::new();
4179            if let Some(b) = &mut self.dsv4 {
4180                let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
4181                crate::dsv4::forward_chunk(
4182                    g,
4183                    layers,
4184                    &cfg,
4185                    st,
4186                    &ids,
4187                    pos,
4188                    &self.inv_freq,
4189                    self.pool.as_deref(),
4190                    &mut lg,
4191                    end == input_ids.len(),
4192                );
4193            }
4194            if end == input_ids.len() {
4195                self.graph_logits = Some(lg);
4196            }
4197            pos = end;
4198            hidden = vec![0.0; self.hidden_size];
4199        }
4200        let dsv41_prefill = self.dsv41_prefill.take();
4201        while self.dsv41.is_some()
4202            && mtp.is_none()
4203            && pos < input_ids.len()
4204            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4205        {
4206            let end = (pos + prefill_chunk()).min(input_ids.len());
4207            let ids: Vec<u32> = input_ids[pos..end].to_vec();
4208            let mut lg = Vec::new();
4209            if let Some(b) = &mut self.dsv41 {
4210                let (g, layers, cfg, st) = (&b.0, &b.1, &b.2, &mut b.3);
4211                if let Some((embeddings, participates)) = dsv41_prefill.as_ref() {
4212                    crate::dsv41::forward_chunk_masked_with_embeddings(
4213                        g,
4214                        layers,
4215                        cfg,
4216                        st,
4217                        &ids,
4218                        pos,
4219                        &embeddings[pos..end],
4220                        &participates[pos..end],
4221                        self.pool.as_deref(),
4222                        &mut lg,
4223                    );
4224                } else {
4225                    crate::dsv41::forward_chunk(
4226                        g,
4227                        layers,
4228                        cfg,
4229                        st,
4230                        &ids,
4231                        pos,
4232                        self.pool.as_deref(),
4233                        &mut lg,
4234                    );
4235                }
4236            }
4237            if end == input_ids.len() {
4238                self.graph_logits = Some(lg);
4239            }
4240            pos = end;
4241            hidden = vec![0.0; self.hidden_size];
4242        }
4243        // With dynamic routing, prefill sequentially so the φ hook fires
4244        // over the PROMPT — the router enters decode with a warm φ (the
4245        // fused-pair path skips the per-layer φ capture). o1 layers
4246        // collect their query trace in both the single and pair paths.
4247        let dyn_prefill = router.is_some();
4248        // Optional bounded calibration prefix for generation.  The normal
4249        // O(1) path seals after the full prompt; this explicit knob instead
4250        // runs only the requested prefix through exact attention, seals the
4251        // Nyström state, and streams the rest of the prompt through the same
4252        // O(1) step used by decode.  It keeps the O(1) layers' Q trace and
4253        // temporary full KV bounded by the prefix while leaving the default
4254        // full-prompt quality profile untouched.
4255        let o1_prefill_limit = o1_prefill
4256            .and_then(|requested| self.o1_effective_boundary(requested))
4257            .map(|boundary| boundary.min(input_ids.len()));
4258        let mut o1_sealed = false;
4259        if let Some(limit) = o1_prefill_limit {
4260            // Reuse the exact batched prefix machinery when available; it
4261            // records the same per-position Q trace as the full prefill.
4262            if self.can_prefill_batched() && limit > 2 {
4263                let chunk = self.prefill_chunk();
4264                let hs = self.hidden_size;
4265                while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4266                    let end = (pos + chunk).min(limit);
4267                    let hb = self.prefill_batch(&input_ids[pos..end], pos);
4268                    hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4269                    pos = end;
4270                }
4271            } else {
4272                while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4273                    hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, None);
4274                    pos += 1;
4275                }
4276            }
4277            if pos >= limit {
4278                o1_sealed = match self.o1_seal_checked() {
4279                    Ok(sealed) => sealed,
4280                    Err(err) => {
4281                        self.finish_generation(&mut mtp, &mut router, true);
4282                        return Err(err);
4283                    }
4284                };
4285                tracing::info!(
4286                    "o1 bounded prompt prefix: requested={} effective={} processed={} of {} token(s)",
4287                    o1_prefill.unwrap_or(0),
4288                    self.o1_effective_boundary(o1_prefill.unwrap_or(0))
4289                        .unwrap_or(limit),
4290                    limit,
4291                    input_ids.len()
4292                );
4293            }
4294        }
4295        // q1 hybrids on Metal: the per-position GPU token graph beats
4296        // the CPU chunk-GEMM (whose wall is the sequential scalar GDN
4297        // recurrence), so prefill goes position-by-position through the
4298        // same graph as decode. Pure-attention models keep the batched
4299        // path — there the chunk-GEMM amortization wins.
4300        let graph_prefill = self.graph_prefill_preferred();
4301        // Native Metal, q4tp GDN hybrids: the prompt through the b-row
4302        // rows graph — projections as GEMMs over up to 512 positions, the
4303        // GDN recurrence in registers on the device, K/V rows appended by
4304        // the chunk — instead of one token-graph submit per position (the
4305        // 27B: 8 tok/s → GEMM-bound). The MTP warm-up rows come out of one
4306        // batched run of the block per chunk. Any refusal leaves the rest
4307        // of the prompt to the sequential paths below.
4308        #[cfg(target_os = "macos")]
4309        if task_mask.is_none()
4310            && !dyn_prefill
4311            && (crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in())
4312            && crate::gpu::enabled_here()
4313            && self.gdn_cfg.is_some()
4314            && self.g3n.is_none()
4315            && input_ids.len() > 8
4316            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
4317            && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
4318        {
4319            let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
4320                .ok()
4321                .and_then(|v| v.parse().ok())
4322                .filter(|&v| (16..=512).contains(&v))
4323                .unwrap_or(256);
4324            let hs = self.hidden_size;
4325            let _tp = std::time::Instant::now();
4326            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4327                let end = (pos + chunk).min(input_ids.len());
4328                let hb = match self.prefill_batch_metal(&input_ids[pos..end], pos) {
4329                    MetalPrefillOutcome::Completed(hb) => hb,
4330                    MetalPrefillOutcome::Declined => break,
4331                    MetalPrefillOutcome::Failed => {
4332                        self.finish_generation(&mut mtp, &mut router, true);
4333                        return Err("ordinary Metal prefill failed after admission".into());
4334                    }
4335                };
4336                if let Some(m) = &mut mtp {
4337                    let n_pairs = if end < input_ids.len() {
4338                        end - pos
4339                    } else {
4340                        end - pos - 1
4341                    };
4342                    if n_pairs > 0 {
4343                        let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
4344                            .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
4345                            .collect();
4346                        if !self.mtp_warm_batch_metal(m, &pairs, pos) {
4347                            for (j, (h, t)) in pairs.iter().enumerate() {
4348                                let h = h.to_vec();
4349                                let _ = self.mtp_step(m, &h, *t, pos + j);
4350                            }
4351                        }
4352                    }
4353                }
4354                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4355                pos = end;
4356            }
4357            if std::env::var("CMF_PREFILL_PROF").is_ok() {
4358                eprintln!(
4359                    "metal-prefill: {} of {} tokens in {:.1} ms",
4360                    pos,
4361                    input_ids.len(),
4362                    _tp.elapsed().as_secs_f64() * 1e3
4363                );
4364            }
4365        }
4366        self.mimo_moe_prepare();
4367        // A MoE stack larger than the card (MiMo-V2 q4tp on 96 GB): the
4368        // batched wgpu graph runs the device prefix of every chunk — its
4369        // experts resident — and the host's batched layer walk finishes
4370        // the chunk. Any refusal leaves the rest of the prompt to the
4371        // chunked prefill below.
4372        #[cfg(not(target_os = "macos"))]
4373        if task_mask.is_none()
4374            && !dyn_prefill
4375            && !graph_prefill
4376            && mtp.is_none()
4377            && o1_prefill.is_none()
4378            && !self.o1_active()
4379            && input_ids.len() > 2
4380            && self.batch_prefix_prefill()
4381        {
4382            let chunk = self.prefill_chunk().max(1);
4383            let hs = self.hidden_size;
4384            let t_bp = std::time::Instant::now();
4385            let pos0 = pos;
4386            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4387                let end = (pos + chunk).min(input_ids.len());
4388                let bk = end - pos;
4389                let mut hiddens = vec![0f32; bk * hs];
4390                for (j, &id) in input_ids[pos..end].iter().enumerate() {
4391                    hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4392                }
4393                let positions: Vec<usize> = (pos..end).collect();
4394                let mut run = 0usize;
4395                let outcome = self.try_batch_graph_wgpu_prefix(
4396                    &mut hiddens,
4397                    &positions,
4398                    bk,
4399                    None,
4400                    Some(&mut run),
4401                );
4402                match outcome {
4403                    crate::gpu::BatchGraphOutcome::Completed => {
4404                        let hb = if run < self.num_layers {
4405                            self.prefill_batch_span(
4406                                PrefillIn::Hidden(&hiddens),
4407                                pos,
4408                                None,
4409                                run,
4410                                self.num_layers,
4411                            )
4412                        } else {
4413                            hiddens
4414                        };
4415                        if mimo_spec {
4416                            self.mimo_note_rows(&hb, pos);
4417                        }
4418                        hidden.copy_from_slice(&hb[(bk - 1) * hs..]);
4419                        pos = end;
4420                    }
4421                    crate::gpu::BatchGraphOutcome::Failed => {
4422                        self.finish_generation(&mut mtp, &mut router, true);
4423                        return Err("batched prefix prefill failed after admission".into());
4424                    }
4425                    crate::gpu::BatchGraphOutcome::Declined => {
4426                        // Earlier chunks left their prefix rows on the
4427                        // device only: the host walk below needs them.
4428                        #[cfg(feature = "gpu")]
4429                        if pos > pos0 {
4430                            self.pull_lagging_host_kv(0, self.num_layers, pos);
4431                        }
4432                        break;
4433                    }
4434                }
4435            }
4436            if std::env::var("CMF_PREFILL_PROF").is_ok() {
4437                eprintln!(
4438                    "batch-prefix prefill: {} of {} tokens in {:.1} ms",
4439                    pos - pos0,
4440                    input_ids.len(),
4441                    t_bp.elapsed().as_secs_f64() * 1e3
4442                );
4443            }
4444        }
4445        if task_mask.is_none()
4446            && !dyn_prefill
4447            && !graph_prefill
4448            && self.can_prefill_batched()
4449            && self.g3n.is_none()
4450            && o1_prefill.is_none()
4451            && input_ids.len() > 2
4452        {
4453            // Production prefill = the same chunked prefill-GEMM that
4454            // bench/PPL measure (roadmap §3 P0: generation used to warm
4455            // the prompt with the slower pair path — the published
4456            // prefill number didn't match real TTFT). MTP warm-up reads
4457            // each position's hidden straight from the chunk result.
4458            let chunk = self.prefill_chunk();
4459            let hs = self.hidden_size;
4460            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4461                let end = (pos + chunk).min(input_ids.len());
4462                let hb = self.prefill_batch(&input_ids[pos..end], pos);
4463                if mimo_spec {
4464                    self.mimo_note_rows(&hb, pos);
4465                }
4466                if let Some(m) = &mut mtp {
4467                    let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4468                        .ok()
4469                        .and_then(|v| v.parse().ok())
4470                        .unwrap_or(0);
4471                    for p in pos..end {
4472                        if p + 1 < input_ids.len() {
4473                            if probe >= 1 && p + 2 < input_ids.len() {
4474                                // Teacher-forced chain acceptance (see the
4475                                // tail loop's twin): the warm-up row stays,
4476                                // the chain's rows roll back.
4477                                let (d1, mut hx) = self.mtp_step_h(
4478                                    m,
4479                                    &hb[(p - pos) * hs..(p - pos + 1) * hs],
4480                                    input_ids[p + 1],
4481                                    p,
4482                                );
4483                                let mut ok = d1 == input_ids[p + 2];
4484                                Self::chain_probe_note(0, ok);
4485                                let mut d_prev = d1;
4486                                let mut extra = 0usize;
4487                                for j in 1..probe {
4488                                    if p + 2 + j >= input_ids.len() {
4489                                        break;
4490                                    }
4491                                    let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
4492                                    extra += 1;
4493                                    ok = ok && dj == input_ids[p + 2 + j];
4494                                    Self::chain_probe_note(j, ok);
4495                                    d_prev = dj;
4496                                    hx = hj;
4497                                }
4498                                m.kv.truncate_last(extra);
4499                            } else {
4500                                let _ = self.mtp_step(
4501                                    m,
4502                                    &hb[(p - pos) * hs..(p - pos + 1) * hs],
4503                                    input_ids[p + 1],
4504                                    p,
4505                                );
4506                            }
4507                        }
4508                    }
4509                }
4510                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4511                pos = end;
4512            }
4513        }
4514        let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
4515        if task_mask.is_none()
4516            && !dyn_prefill
4517            && !graph_prefill
4518            && !pair_off
4519            && self.pair_supported()
4520            && o1_prefill.is_none()
4521        {
4522            while pos + 1 < input_ids.len()
4523                && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4524            {
4525                let e1 = self.embed_single(input_ids[pos]);
4526                let e2 = self.embed_single(input_ids[pos + 1]);
4527                let (h1, h2) = self.forward_pair(&e1, &e2, pos);
4528                if mimo_spec {
4529                    self.mimo_note_rows(&h1, pos);
4530                    self.mimo_note_rows(&h2, pos + 1);
4531                }
4532                // Both prefill tokens are real → commit lane-2 states.
4533                self.commit_linear_scratch();
4534                if let Some(m) = &mut mtp {
4535                    let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
4536                    if pos + 2 < input_ids.len() {
4537                        let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4538                            .ok()
4539                            .and_then(|v| v.parse().ok())
4540                            .unwrap_or(0);
4541                        if probe >= 1 && pos + 3 < input_ids.len() {
4542                            // Same teacher-forced chain table as the tail
4543                            // loop below, fed from the pair path that owns
4544                            // most prefill positions.
4545                            let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
4546                            let mut ok = d1 == input_ids[pos + 3];
4547                            Self::chain_probe_note(0, ok);
4548                            let mut d_prev = d1;
4549                            let mut extra = 0usize;
4550                            for j in 1..probe {
4551                                if pos + 3 + j >= input_ids.len() {
4552                                    break;
4553                                }
4554                                let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
4555                                extra += 1;
4556                                ok = ok && dj == input_ids[pos + 3 + j];
4557                                Self::chain_probe_note(j, ok);
4558                                d_prev = dj;
4559                                hx = hj;
4560                            }
4561                            m.kv.truncate_last(extra);
4562                        } else {
4563                            let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
4564                        }
4565                    }
4566                }
4567                hidden = h2;
4568                pos += 2;
4569            }
4570        }
4571        // Batched GPU prefill for the wgpu decode graph (GDN hybrids): K prompt
4572        // positions per submit — projections/FFN as GEMMs (weight once per K),
4573        // attention/GDN looped inside — instead of one whole-graph submit per
4574        // position. Falls through to the per-position graph on any refusal.
4575        // Batched prefill is opt-in (CMF_BATCH_K>0). Default 0 = per-position
4576        // graph prefill. (Steady-state decode is provably identical either way —
4577        // token-graph submit and lm_head both unchanged — so this only trades
4578        // prefill wall.)
4579        // A bounded O(1) prefix is the one post-seal prompt interval: only
4580        // admit its batch when the device O(1) route is explicitly enabled and
4581        // every sealed layer exposes a portable view. The same batch size and
4582        // refusal behavior remain the ordinary controls/comparator.
4583        let o1_batch_ready = o1_sealed
4584            && o1_prefill.is_some()
4585            && mtp.is_none()
4586            && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
4587            && (0..self.num_layers).all(|li| {
4588                let cache = &self.kv_cache.layers[self.phys_layer(li)];
4589                cache.o1.is_none() || cache.o1_views().is_some()
4590            });
4591        // The ordinary graph-prefill route can share each completed trunk
4592        // chunk with an attached MTP head.  Keep chain probing on its
4593        // established per-position path: the probe deliberately needs every
4594        // teacher-forced draft row and its rollback table.
4595        let mtp_batch_prefill = mtp.is_some()
4596            && graph_prefill
4597            && task_mask.is_none()
4598            && !dyn_prefill
4599            && !self.o1_active()
4600            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err();
4601        if batch_k > 0
4602            && (graph_prefill || o1_batch_ready)
4603            && task_mask.is_none()
4604            && (!self.o1_active() || o1_batch_ready)
4605            && (mtp.is_none() || mtp_batch_prefill)
4606            && !dyn_prefill
4607            && pos + 1 < input_ids.len()
4608        {
4609            let hs = self.hidden_size;
4610            let chunk = batch_k;
4611            while pos < input_ids.len() {
4612                let end = (pos + chunk).min(input_ids.len());
4613                let bk = end - pos;
4614                let mut hiddens = vec![0f32; bk * hs];
4615                for (j, &id) in input_ids[pos..end].iter().enumerate() {
4616                    hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4617                }
4618                let positions: Vec<usize> = (pos..end).collect();
4619                let t_chunk = std::time::Instant::now();
4620                let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
4621                let ok_b = outcome == crate::gpu::BatchGraphOutcome::Completed;
4622                if std::env::var("CMF_GRAPH_PROF").is_ok() {
4623                    let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
4624                    eprintln!(
4625                        "batch-chunk: phase=prompt mode={} k={bk} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
4626                        if o1_batch_ready {
4627                            "o1"
4628                        } else if mtp_batch_prefill {
4629                            "ordinary_mtp"
4630                        } else {
4631                            "ordinary"
4632                        },
4633                        bk as f64 / (ms / 1000.0)
4634                    );
4635                }
4636                {
4637                    use std::sync::atomic::{AtomicBool, Ordering};
4638                    static SAID: AtomicBool = AtomicBool::new(false);
4639                    if !SAID.swap(true, Ordering::Relaxed) {
4640                        if ok_b {
4641                            tracing::info!(
4642                                "batched prefill: ACTIVE mode={} (k={bk})",
4643                                if o1_batch_ready {
4644                                    "o1"
4645                                } else if mtp_batch_prefill {
4646                                    "ordinary_mtp"
4647                                } else {
4648                                    "ordinary"
4649                                }
4650                            );
4651                        } else {
4652                            tracing::warn!("batched prefill {:?} — per-position graph", outcome);
4653                        }
4654                    }
4655                }
4656                if ok_b {
4657                    if mimo_spec {
4658                        self.mimo_note_rows(&hiddens, pos);
4659                    }
4660                    if mtp_batch_prefill {
4661                        let n_pairs = mtp_prefill_pair_count(pos, end, input_ids.len());
4662                        if n_pairs > 0 {
4663                            // `hiddens` is owned by this chunk, so materialize
4664                            // row slices before borrowing the detached MTP
4665                            // module.  The last prompt row has no successor;
4666                            // the helper above is the single source of that
4667                            // boundary rule.
4668                            let rows: Vec<Vec<f32>> = (0..n_pairs)
4669                                .map(|j| hiddens[j * hs..(j + 1) * hs].to_vec())
4670                                .collect();
4671                            let pairs: Vec<(&[f32], u32)> = rows
4672                                .iter()
4673                                .enumerate()
4674                                .map(|(j, row)| (row.as_slice(), input_ids[pos + j + 1]))
4675                                .collect();
4676                            if std::env::var("CMF_GRAPH_PROF").is_ok() {
4677                                eprintln!(
4678                                    "mtp-warm: phase=prompt mode=ordinary_mtp first_pos={} pairs={} last_pos={}",
4679                                    pos,
4680                                    n_pairs,
4681                                    pos + n_pairs - 1,
4682                                );
4683                            }
4684                            let warm_error = if let Some(m) = mtp.as_mut() {
4685                                self.mtp_warm_prefill_pairs(m, &pairs, pos).err()
4686                            } else {
4687                                None
4688                            };
4689                            if let Some(err) = warm_error {
4690                                // The trunk batch was already admitted.  A
4691                                // failed MTP warm-up therefore clears both
4692                                // mirrors and exits; continuing would pair a
4693                                // current trunk state with a stale MTP cache.
4694                                self.finish_generation(&mut mtp, &mut router, true);
4695                                return Err(err.to_string());
4696                            }
4697                        }
4698                    }
4699                    hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
4700                    pos = end;
4701                } else if outcome == crate::gpu::BatchGraphOutcome::Failed {
4702                    // A failed batch may have advanced a device recurrent
4703                    // state (ordinary GDN or sealed O(1)). A CPU fallback
4704                    // would then observe stale accumulators, so clear the
4705                    // request state and make the failure explicit.
4706                    self.finish_generation(&mut mtp, &mut router, true);
4707                    return Err(if o1_batch_ready {
4708                        "sealed O(1) batch graph failed after admission".to_string()
4709                    } else {
4710                        "ordinary recurrent batch graph failed after admission".to_string()
4711                    });
4712                } else {
4713                    break; // unsupported → per-position graph handles the rest
4714                }
4715            }
4716        }
4717        // Resident Embryo graph: the prompt in chunks of one submit each
4718        // instead of one whole-graph submit per position; the last chunk
4719        // carries the logits exactly as the per-position walk would.
4720        if graph_prefill
4721            && task_mask.is_none()
4722            && mtp.is_none()
4723            && !dyn_prefill
4724            && pos == 0
4725            && input_ids.len() > 1
4726            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4727        {
4728            if let Some(lg) = self.embryo_prefill_chunked(input_ids, 0) {
4729                self.graph_logits = Some(lg);
4730                hidden = vec![0.0; self.hidden_size];
4731                pos = input_ids.len();
4732            }
4733        }
4734        while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4735            self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
4736            hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
4737            if mimo_spec {
4738                self.mimo_note_rows(&hidden, pos);
4739            }
4740            if let Some(m) = &mut mtp {
4741                if pos + 1 < input_ids.len() {
4742                    // `CMF_MTP_CHAIN_PROBE=k`: teacher-forced acceptance of a
4743                    // CHAINED draft — iterate the head on its own hidden k
4744                    // deep and score every depth against the prompt's real
4745                    // continuation. The economics of a k-token speculative
4746                    // round stand or fall on this table.
4747                    let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4748                        .ok()
4749                        .and_then(|v| v.parse().ok())
4750                        .unwrap_or(0);
4751                    if probe >= 1 && pos + 2 < input_ids.len() {
4752                        let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
4753                        let mut ok = d1 == input_ids[pos + 2];
4754                        Self::chain_probe_note(0, ok);
4755                        let mut d_prev = d1;
4756                        let mut extra = 0usize;
4757                        for j in 1..probe {
4758                            if pos + 2 + j >= input_ids.len() {
4759                                break;
4760                            }
4761                            let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
4762                            extra += 1;
4763                            ok = ok && dj == input_ids[pos + 2 + j];
4764                            Self::chain_probe_note(j, ok);
4765                            d_prev = dj;
4766                            hx = hj;
4767                        }
4768                        // The chain's rows are speculation, not the prompt —
4769                        // keep only the warmup row the plain path would add.
4770                        m.kv.truncate_last(extra);
4771                    } else {
4772                        let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
4773                    }
4774                }
4775            }
4776            pos += 1;
4777        }
4778        if std::env::var("CMF_PREFILL_PROF").is_ok() {
4779            eprintln!(
4780                "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
4781                input_ids.len(),
4782                _tpf.elapsed().as_secs_f64() * 1000.0
4783            );
4784        }
4785        if self
4786            .graph_failed
4787            .swap(false, std::sync::atomic::Ordering::Relaxed)
4788        {
4789            // MTP is detached for speculative generation.  Restore the
4790            // module before returning the terminal graph error; otherwise a
4791            // failed request would silently remove the head from a pooled
4792            // pipeline and the next request would lose its configured route.
4793            self.finish_generation(&mut mtp, &mut router, true);
4794            return Err("GPU token graph failed during prefill".to_string());
4795        }
4796        // Cancelled mid-prefill: the cache holds a partial prompt —
4797        // drop the reuse history and return an empty generation.
4798        if self
4799            .cancel
4800            .swap(false, std::sync::atomic::Ordering::Relaxed)
4801        {
4802            // A cancelled prefill can already have advanced the device
4803            // mirror. Drop the whole partial sequence so a pooled pipeline
4804            // cannot carry that state into its next request.
4805            self.finish_generation(&mut mtp, &mut router, true);
4806            return Ok(GenerateResult {
4807                text: String::new(),
4808                token_ids: Vec::new(),
4809                prompt_tokens: input_ids.len(),
4810                tokens_generated: 0,
4811                finish_reason: "cancelled".to_string(),
4812                mtp_drafted: 0,
4813                mtp_accepted: 0,
4814                token_confidence: Vec::new(),
4815                traces: Vec::new(),
4816            });
4817        }
4818
4819        // Prompt absorbed → freeze the o1 layers' skeletons; from here
4820        // every decode step on those layers is O(W + m·dv + m²).
4821        if !o1_sealed {
4822            match self.o1_seal_checked() {
4823                Ok(_) => {}
4824                Err(err) => {
4825                    self.finish_generation(&mut mtp, &mut router, true);
4826                    return Err(err);
4827                }
4828            }
4829        }
4830
4831        // Commit one token: push, check EOS, stream. Returns false = stop.
4832        macro_rules! commit {
4833            ($id:expr) => {{
4834                all_ids.push($id);
4835                generated += 1;
4836                self.note_draft_id($id);
4837                if self.tokenizer.is_eos($id) && !self.ignore_eos {
4838                    finish_reason = "stop".to_string();
4839                    false
4840                } else {
4841                    let token_text = self.tokenizer.decode_token($id);
4842                    let mut go = true;
4843                    if let Some(ref mut cb) = on_token {
4844                        if !cb(&token_text) {
4845                            finish_reason = "cancelled".to_string();
4846                            go = false;
4847                        }
4848                    }
4849                    go
4850                }
4851            }};
4852        }
4853
4854        // Speculation is decided by MEASUREMENT, not by an acceptance
4855        // model. A k=4 round costs ~3.8 plain tokens on the 5090 (draft
4856        // 6.6 + verify 66.6 + commit 4.8 ms against a 20.6 ms token), so it
4857        // pays only when the head lands ~2.8 of 4 — predictable text (code,
4858        // structured output) does, free prose often does not, and the
4859        // ratio at which the two cross depends on the card and the context
4860        // depth. So: four speculative rounds timed, then eight plain
4861        // tokens timed, and the faster arm runs until a re-check 256
4862        // tokens later (context growth moves the balance). The trial
4863        // costs at most a few tokens of the slower arm per 256.
4864        let mut spec_trial = SpecTrial::Spec {
4865            t0: std::time::Instant::now(),
4866            gen0: generated,
4867            rounds: 0,
4868        };
4869        // The token-count proxy prices a round at ~1.9 plain tokens. That
4870        // holds for the Metal rounds whose cost was measured — greedy and
4871        // the sparse sampling chain — so an expensive round (the dense
4872        // chain, reachable only by `CMF_GRAPH_SPEC_SAMPLE=1`) still times
4873        // the plain path before it decides.
4874        let mut spec_mon = SpecMon {
4875            metal: graph_spec && crate::gpu::q1_force() && spec_cheap_round,
4876            ..SpecMon::default()
4877        };
4878        let mut spec_watchdog_off = false;
4879        // CMF_GRAPH_SPEC_TIME: the round walls so far (round 1 excluded —
4880        // it pays the scratch), for the outlier test on each new one
4881        let mut spec_walls: Vec<f32> = Vec::new();
4882        // ... and the end of the last round: the host time between rounds
4883        // (token commits, streaming, the loop top) is printed at level 2
4884        let mut spec_round_end: Option<std::time::Instant> = None;
4885        if mimo_spec {
4886            if let Ok(path) = std::env::var("CMF_MIMO_MTP_PROBE") {
4887                if let Some(mut st) = self.mimo_mtp.take() {
4888                    self.mimo_mtp_probe(&mut st, input_ids, &path);
4889                    self.mimo_mtp = Some(st);
4890                }
4891            }
4892        }
4893        // ── Decode ──
4894        let mut next_pos = input_ids.len();
4895        'decode: while generated < max_tokens {
4896            if self
4897                .graph_failed
4898                .swap(false, std::sync::atomic::Ordering::Relaxed)
4899            {
4900                // Keep the detached MTP module attached after a terminal
4901                // graph error so the pipeline can be reused for a fresh
4902                // sequence.  `clear_sequence_state` only clears mirrors and
4903                // host KV; it cannot recover a module dropped here.
4904                self.finish_generation(&mut mtp, &mut router, true);
4905                return Err("GPU token graph failed during decode".to_string());
4906            }
4907            if self
4908                .cancel
4909                .swap(false, std::sync::atomic::Ordering::Relaxed)
4910            {
4911                finish_reason = "cancelled".to_string();
4912                break 'decode;
4913            }
4914            // A rejected speculative draft already drew this position's
4915            // token from the residual distribution (graph_spec_step); it
4916            // is committed as-is — sampling again from the row's logits
4917            // would bias the stream toward the target's mode.
4918            if mimo_spec && next_pos > 0 {
4919                // Every path leaves `hidden` = the backbone output at
4920                // next_pos-1; the draft layers read it (idempotent).
4921                self.mimo_note_rows(&hidden, next_pos - 1);
4922            }
4923            let forced = self.spec_forced.take();
4924            let mut logits = match (forced, self.graph_logits.take()) {
4925                (Some(_), _) => Vec::new(),
4926                (None, Some(lg)) => lg,
4927                (None, None) => {
4928                    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Head);
4929                    inference::rms_norm_into(
4930                        &hidden,
4931                        &self.weights.final_norm,
4932                        self.rms_eps,
4933                        self.norm_style,
4934                        &mut self.ws.n1,
4935                    );
4936                    self.lm_head_forward(&self.ws.n1)
4937                }
4938            };
4939            // CMF_LOGIT_DUMP=<path>: the first decode step's hidden + logits
4940            // as raw f32 (hidden first) — cross-backend numerics diffing.
4941            if generated
4942                == std::env::var("CMF_LOGIT_DUMP_STEP")
4943                    .ok()
4944                    .and_then(|v| v.parse().ok())
4945                    .unwrap_or(0)
4946            {
4947                if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
4948                    let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
4949                    for v in hidden.iter().chain(logits.iter()) {
4950                        bytes.extend_from_slice(&v.to_le_bytes());
4951                    }
4952                    if let Err(e) = std::fs::write(&path, &bytes) {
4953                        eprintln!("logit dump: failed to write {path}: {e}");
4954                        self.finish_generation(&mut mtp, &mut router, true);
4955                        return Err(format!("logit dump write failed: {e}"));
4956                    }
4957                }
4958            }
4959            // CMF_LOGIT_DUMP_ALL=<dir>: every decode step's logits as raw
4960            // f32, `<dir>/step{n:05}.f32` — step-by-step backend diffing
4961            // (a greedy run on two backends compares until they diverge).
4962            if let Ok(dir) = std::env::var("CMF_LOGIT_DUMP_ALL") {
4963                if !logits.is_empty() {
4964                    let path = std::path::Path::new(&dir).join(format!("step{generated:05}.f32"));
4965                    let bytes: Vec<u8> = logits.iter().flat_map(|v| v.to_le_bytes()).collect();
4966                    if let Err(e) =
4967                        std::fs::create_dir_all(&dir).and_then(|_| std::fs::write(&path, &bytes))
4968                    {
4969                        eprintln!("logit dump: failed to write {}: {e}", path.display());
4970                    }
4971                }
4972            }
4973            let t_next = match forced {
4974                Some(c) => c,
4975                None => {
4976                    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Sampler);
4977                    sampler::sample_with_scratch_pool(
4978                        &logits,
4979                        &self.sampler_config,
4980                        self.sampler_config.penalty_past(&all_ids, bounded_native),
4981                        &mut self.rng,
4982                        &mut self.sampler_scratch,
4983                        self.pool.as_deref(),
4984                    )
4985                }
4986            };
4987            if self.confidence_on {
4988                confidence.push(if logits.is_empty() {
4989                    0.0
4990                } else {
4991                    sampler::top1_prob_pool(
4992                        self.pool.as_deref(),
4993                        &mut self.sampler_scratch,
4994                        &logits,
4995                        t_next,
4996                        calib_temp,
4997                    )
4998                });
4999            }
5000            if !logits.is_empty() {
5001                attention::recycle_buf(&mut logits);
5002            }
5003            if trace_on {
5004                // active_skill = the overlay in force while this token was
5005                // generated; recon/switched are filled after the post-emit
5006                // routing eval below (freshest coherence for this token).
5007                let skill = router.as_ref().and_then(|r| r.active_id());
5008                traces.push(TokenTrace {
5009                    t: generated,
5010                    token_id: t_next,
5011                    confidence: confidence.last().copied().unwrap_or(0.0),
5012                    active_skill: skill,
5013                    recon: None,
5014                    switched: false,
5015                });
5016            }
5017            if !commit!(t_next) {
5018                break 'decode;
5019            }
5020            if generated >= max_tokens {
5021                break 'decode;
5022            }
5023
5024            if self.dsv41.is_none() && self.kv_cache.needs_eviction() {
5025                // Say it ONCE, loudly: past this point the model keeps
5026                // talking but has lost half its context, and on a GDN
5027                // hybrid the graph's device state goes stale on top. The
5028                // Qwen3.8 bring-up spent a day reading this cliff as
5029                // three different model bugs.
5030                static SAID: std::sync::Once = std::sync::Once::new();
5031                SAID.call_once(|| {
5032                    tracing::warn!(
5033                        "KV cache full at {} positions — evicting half; quality \
5034                         will degrade. Raise CMF_MAX_SEQ.",
5035                        self.kv_cache.max_seq_len,
5036                    );
5037                });
5038                let keep = (self.kv_cache.max_seq_len / 2).max(1);
5039                self.kv_cache.evict(keep);
5040            }
5041
5042            // Advance the speculation trial: plain-phase accounting and
5043            // the periodic re-check happen here, on every token.
5044            if graph_spec {
5045                match spec_trial {
5046                    SpecTrial::Plain { t0, gen0 } if spec_mon.plain_done(t0, gen0, generated) => {
5047                        spec_mon.plain_ms =
5048                            t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
5049                        let keep = spec_mon.pays();
5050                        tracing::info!(
5051                            "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5052                            spec_mon.tokens,
5053                            spec_mon.round_ms,
5054                            spec_mon.plain_ms,
5055                            if keep { "speculating" } else { "plain" }
5056                        );
5057                        spec_mon.fails = 0;
5058                        spec_trial = SpecTrial::Decided {
5059                            spec: keep,
5060                            recheck_at: if keep { usize::MAX } else { generated + 128 },
5061                        };
5062                    }
5063                    SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
5064                        spec_mon.n = 0;
5065                        spec_trial = SpecTrial::Spec {
5066                            t0: std::time::Instant::now(),
5067                            gen0: generated,
5068                            rounds: 0,
5069                        };
5070                    }
5071                    _ => {}
5072                }
5073                spec_watchdog_off = matches!(
5074                    spec_trial,
5075                    SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
5076                );
5077            }
5078            // ── MiMo-V2 draft stack: draft K, verify K+1 rows in one batch ──
5079            if mimo_spec && generated + 1 < max_tokens && next_pos > 0 {
5080                let budget = max_tokens - generated - 1;
5081                if let Some(mut st) = self.mimo_mtp.take() {
5082                    let k = st.depth.min(budget);
5083                    let r = self.mimo_spec_round(&mut st, next_pos, &all_ids, k);
5084                    self.mimo_mtp = Some(st);
5085                    let r = match r {
5086                        Ok(r) => r,
5087                        Err(err) => {
5088                            self.finish_generation(&mut mtp, &mut router, true);
5089                            return Err(err);
5090                        }
5091                    };
5092                    if let Some(r) = r {
5093                        drafted += r.drafted;
5094                        accepted += r.accepted.len();
5095                        let mut stopped = false;
5096                        for &id in &r.accepted {
5097                            if self.confidence_on {
5098                                confidence.push(0.0);
5099                            }
5100                            if !commit!(id) {
5101                                stopped = true;
5102                                break;
5103                            }
5104                        }
5105                        if stopped {
5106                            break 'decode;
5107                        }
5108                        next_pos += r.accepted.len() + 1;
5109                        hidden = r.hidden;
5110                        // The loop top chooses the round's own token from
5111                        // these logits — the same sampler, same history.
5112                        self.graph_logits = Some(r.logits);
5113                        continue 'decode;
5114                    }
5115                }
5116            }
5117            // ── Qwen3.8-Flash-Next draft head: k greedy drafts from the MTP
5118            //    sidecar, one batched verify window on the device ──
5119            #[cfg(feature = "gpu")]
5120            if self.speculative
5121                && self.qwen4_exp.is_some()
5122                && task_mask.is_none()
5123                && self.sampler_config.temperature < 1e-6
5124                && generated + 1 < max_tokens
5125                && next_pos > 0
5126                && std::env::var("CMF_QWEN_MTP").as_deref() != Ok("0")
5127            {
5128                let r = match &mut self.qwen4_exp {
5129                    Some(b) => crate::qwen4_exp::spec_round(
5130                        &b.0,
5131                        &b.1,
5132                        &b.2,
5133                        &mut b.3,
5134                        next_pos,
5135                        &all_ids,
5136                        &self.inv_freq,
5137                        self.pool.as_deref(),
5138                    ),
5139                    None => None,
5140                };
5141                if let Some(r) = r {
5142                    drafted += r.drafted;
5143                    accepted += r.accepted.len();
5144                    let mut stopped = false;
5145                    for &id in &r.accepted {
5146                        if self.confidence_on {
5147                            confidence.push(0.0);
5148                        }
5149                        if !commit!(id) {
5150                            stopped = true;
5151                            break;
5152                        }
5153                    }
5154                    if stopped {
5155                        break 'decode;
5156                    }
5157                    next_pos += r.accepted.len() + 1;
5158                    hidden.fill(0.0);
5159                    self.graph_logits = Some(r.logits);
5160                    continue 'decode;
5161                }
5162            }
5163            match &mut mtp {
5164                // ── Graph speculation: chain-draft, batch-verify on device ──
5165                #[cfg(feature = "gpu")]
5166                Some(m)
5167                    if graph_spec
5168                        && !spec_watchdog_off
5169                        && generated + 1 < max_tokens
5170                        && next_pos > 0 =>
5171                {
5172                    let t_round = std::time::Instant::now();
5173                    if spec_time_level() >= 2 {
5174                        if let Some(t) = spec_round_end.take() {
5175                            eprintln!(
5176                                "spec-gap {:.2} ms (host between rounds)",
5177                                t.elapsed().as_secs_f64() * 1e3
5178                            );
5179                        }
5180                    }
5181                    spec_stamps_begin();
5182                    // device buffers allocated during this round: a
5183                    // first-touch Shared allocation is zero-filled inside
5184                    // the command buffer that uses it, which is what the
5185                    // long outlier rounds were
5186                    #[cfg(target_os = "macos")]
5187                    let allocs0 = crate::gpu_metal::IO_BUF_ALLOCS
5188                        .load(std::sync::atomic::Ordering::Relaxed);
5189                    #[cfg(not(target_os = "macos"))]
5190                    let allocs0 = 0u64;
5191                    if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
5192                        m,
5193                        &hidden,
5194                        t_next,
5195                        next_pos,
5196                        &mut drafted,
5197                        &mut accepted,
5198                        &mut all_ids,
5199                        max_tokens - generated,
5200                    ) {
5201                        next_pos = n_pos;
5202                        hidden = new_h;
5203                        let level = spec_time_level();
5204                        if level > 0 {
5205                            let wall = t_round.elapsed().as_secs_f32() * 1e3;
5206                            let stamps = spec_stamps_take();
5207                            // the running median of the rounds before this
5208                            // one (round 1 pays the scratch: not a sample)
5209                            let median = if spec_walls.len() >= 3 {
5210                                let mut s = spec_walls.clone();
5211                                s.sort_by(|a, b| a.partial_cmp(b).unwrap());
5212                                Some(s[s.len() / 2])
5213                            } else {
5214                                None
5215                            };
5216                            let outlier = median.is_some_and(|m| wall > 1.4 * m);
5217                            #[cfg(target_os = "macos")]
5218                            let allocs = crate::gpu_metal::IO_BUF_ALLOCS
5219                                .load(std::sync::atomic::Ordering::Relaxed)
5220                                - allocs0;
5221                            #[cfg(not(target_os = "macos"))]
5222                            let allocs = allocs0;
5223                            eprintln!(
5224                                "spec-round wall {wall:.1} ms → {} tokens{}{}",
5225                                extra.len() + 1,
5226                                if allocs > 0 {
5227                                    format!(" [{allocs} new device buffers]")
5228                                } else {
5229                                    String::new()
5230                                },
5231                                match (outlier, median) {
5232                                    (true, Some(m)) => format!(" OUTLIER (median {m:.1})"),
5233                                    _ => String::new(),
5234                                }
5235                            );
5236                            if level >= 2 || outlier {
5237                                let sum: f32 = stamps.iter().map(|s| s.1).sum();
5238                                eprintln!(
5239                                    "spec-stamps: {}| untracked {:.1}",
5240                                    spec_stamps_format(&stamps),
5241                                    wall - sum
5242                                );
5243                            }
5244                            if spec_mon.n >= 1 {
5245                                spec_walls.push(wall);
5246                            }
5247                        }
5248                        // One speculative round done: the monitor counts it
5249                        // (round 1 untimed — it pays the batch scratch and
5250                        // the draft mirror), and the trial advances.
5251                        spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
5252                        // the round's tokens land in `generated` below; the
5253                        // plain phase must start counting AFTER them
5254                        spec_trial = Self::spec_trial_round(
5255                            spec_trial,
5256                            &mut spec_mon,
5257                            generated + extra.len() + 1,
5258                        );
5259                        let mut stopped = false;
5260                        for &id in &extra {
5261                            if self.confidence_on {
5262                                confidence.push(0.0);
5263                            }
5264                            if !commit!(id) {
5265                                stopped = true;
5266                                break;
5267                            }
5268                        }
5269                        if stopped {
5270                            break 'decode;
5271                        }
5272                        if spec_time_level() >= 2 {
5273                            spec_round_end = Some(std::time::Instant::now());
5274                        }
5275                        continue 'decode;
5276                    }
5277                    if self
5278                        .graph_failed
5279                        .swap(false, std::sync::atomic::Ordering::Relaxed)
5280                    {
5281                        // `graph_spec_step` may have detached MTP while a
5282                        // warm-up was in flight.  Do not reinterpret its
5283                        // terminal device failure as a plain decode step;
5284                        // restore the head, clear both mirrors, and surface
5285                        // one explicit error to the caller.
5286                        self.finish_generation(&mut mtp, &mut router, true);
5287                        return Err("GPU MTP graph failed during speculative decode".to_string());
5288                    }
5289                    // Declined (batch graph refused): plain forward below —
5290                    // and a round that produced one token for the trial's
5291                    // ledger, so a graph that keeps refusing is measured out
5292                    // like a head that keeps missing (it was spinning
5293                    // forever on a file whose batch graph declines).
5294                    // A declined round is not a cheap one-token round — it
5295                    // is a verify that does not exist for this file (a
5296                    // healed q8_2f tail measured 760 drafts, 0 accepted, 33
5297                    // against 48.8 tok/s while the monitor called the draft
5298                    // alone "paying"). Count it as the losing streak in one.
5299                    spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
5300                    spec_mon.tokens = 0.0;
5301                    spec_mon.fails = 3;
5302                    spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
5303                    hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
5304                    next_pos += 1;
5305                    continue 'decode;
5306                }
5307                // ── Speculative: draft t+2, verify in a fused pair ──
5308                Some(m) if !graph_spec && generated + 1 < max_tokens => {
5309                    let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
5310                    drafted += 1;
5311                    let emb1 = self.embed_single(t_next);
5312                    let emb2 = self.embed_single(draft);
5313                    let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
5314
5315                    inference::rms_norm_into(
5316                        &h1,
5317                        &self.weights.final_norm,
5318                        self.rms_eps,
5319                        self.norm_style,
5320                        &mut self.ws.n1,
5321                    );
5322                    let mut logits1 = self.lm_head_forward(&self.ws.n1);
5323                    let t_after = sampler::sample_with_scratch_pool(
5324                        &logits1,
5325                        &self.sampler_config,
5326                        self.sampler_config.penalty_past(&all_ids, bounded_native),
5327                        &mut self.rng,
5328                        &mut self.sampler_scratch,
5329                        self.pool.as_deref(),
5330                    );
5331                    if self.confidence_on {
5332                        confidence.push(sampler::top1_prob_pool(
5333                            self.pool.as_deref(),
5334                            &mut self.sampler_scratch,
5335                            &logits1,
5336                            t_after,
5337                            calib_temp,
5338                        ));
5339                    }
5340                    attention::recycle_buf(&mut logits1);
5341                    if trace_on {
5342                        // Speculative decode is mutually exclusive with
5343                        // dynamic routing (router is None here) — no skill.
5344                        traces.push(TokenTrace {
5345                            t: generated,
5346                            token_id: t_after,
5347                            confidence: confidence.last().copied().unwrap_or(0.0),
5348                            active_skill: None,
5349                            recon: None,
5350                            switched: false,
5351                        });
5352                    }
5353                    let stop = !commit!(t_after);
5354
5355                    if t_after == draft {
5356                        accepted += 1;
5357                        self.commit_linear_scratch();
5358                        let _ = self.mtp_step(m, &h1, t_after, next_pos);
5359                        hidden = h2;
5360                        next_pos += 2;
5361                    } else {
5362                        // The draft lane is wrong: roll its KV entry back.
5363                        for layer in &mut self.kv_cache.layers {
5364                            layer.truncate_last(1);
5365                        }
5366                        if !stop {
5367                            let _ = self.mtp_step(m, &h1, t_after, next_pos);
5368                            hidden = self.forward_layers(
5369                                &self.embed_single(t_after),
5370                                next_pos + 1,
5371                                None,
5372                            );
5373                        }
5374                        next_pos += 2;
5375                    }
5376                    if stop {
5377                        break 'decode;
5378                    }
5379                }
5380                // ── Vanilla: forward the sampled token ──
5381                _ => {
5382                    // ── DeepSeek-V4 speculative decode (CMF_DSV4_SPEC=1):
5383                    // draft five on the card, verify batched, commit the
5384                    // accepted prefix. Greedy only; a rejected token's state
5385                    // is restored and replayed, so output equals the walk. ──
5386                    #[cfg(feature = "gpu")]
5387                    if Self::dsv4_spec_on() && self.dsv4.is_some() {
5388                        static SAID: std::sync::Once = std::sync::Once::new();
5389                        SAID.call_once(|| {
5390                            eprintln!(
5391                                "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
5392                                !self.dsv4_mtp.is_empty(),
5393                                task_mask.is_none(),
5394                                router.is_none(),
5395                                !trace_on,
5396                                self.sampler_config.temperature < 1e-6,
5397                                self.sampler_config.repetition_penalty == 1.0,
5398                            );
5399                        });
5400                    }
5401                    #[cfg(feature = "gpu")]
5402                    if Self::dsv4_spec_on()
5403                        && self.dsv4.is_some()
5404                        && !self.dsv4_mtp.is_empty()
5405                        && task_mask.is_none()
5406                        && router.is_none()
5407                        && !trace_on
5408                        && self.sampler_config.temperature < 1e-6
5409                        && self.sampler_config.repetition_penalty == 1.0
5410                        && generated + 1 < max_tokens
5411                        && all_ids.len() >= 2
5412                        && generated >= dsv4_spec_retry_at
5413                    {
5414                        let tip_token = all_ids[all_ids.len() - 2];
5415                        let drafted0 = drafted;
5416                        let round = self.dsv4_spec_step(
5417                            tip_token,
5418                            t_next,
5419                            next_pos,
5420                            max_tokens.saturating_sub(generated),
5421                            &mut drafted,
5422                            &mut accepted,
5423                        );
5424                        if drafted > drafted0 {
5425                            let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
5426                            if useful {
5427                                dsv4_spec_bad = 0;
5428                            } else {
5429                                dsv4_spec_bad += 1;
5430                                if dsv4_spec_bad >= 2 {
5431                                    dsv4_spec_bad = 0;
5432                                    dsv4_spec_retry_at = generated.saturating_add(32);
5433                                    tracing::info!(
5434                                        "dsv4: draft не окупился дважды — точный walk на 32 токена"
5435                                    );
5436                                }
5437                            }
5438                        }
5439                        if let Some((extra, n_pos)) = round {
5440                            next_pos = n_pos;
5441                            let mut stopped = false;
5442                            for &id in &extra {
5443                                if self.confidence_on {
5444                                    confidence.push(0.0);
5445                                }
5446                                if !commit!(id) {
5447                                    stopped = true;
5448                                    break;
5449                                }
5450                            }
5451                            if stopped {
5452                                break 'decode;
5453                            }
5454                            continue 'decode;
5455                        }
5456                    }
5457                    self.graph_want_logits = fuse_lm;
5458                    // Greedy burst (CMF_MULTISTEP, default 8, 1 = off): while
5459                    // nothing observes per-token state — pure argmax sampling,
5460                    // no router/trace/confidence/mask — decode k tokens per
5461                    // submit and commit them wholesale. The trailing normal
5462                    // forward leaves logits for the loop top, as always.
5463                    let mut t_fwd = t_next;
5464                    let pure_greedy = self.sampler_config.temperature < 1e-6
5465                        && self.sampler_config.repetition_penalty == 1.0
5466                        && self.sampler_config.suppress_tokens.is_empty();
5467                    // Off by default: at every k the burst measured at or
5468                    // below the plain path on this graph shape (k=1 loses
5469                    // the argmax dispatches vs a 1 MB readback, k>=8 loses
5470                    // inter-step drains vs the saved sync). Experimental.
5471                    let burst_k = std::env::var("CMF_MULTISTEP")
5472                        .ok()
5473                        .and_then(|v| v.parse::<usize>().ok())
5474                        .unwrap_or(0);
5475                    if pure_greedy
5476                        && burst_k >= 1
5477                        && fuse_lm
5478                        && task_mask.is_none()
5479                        && router.is_none()
5480                        && !trace_on
5481                        && !self.confidence_on
5482                    {
5483                        let mut stopped = false;
5484                        loop {
5485                            let room = max_tokens.saturating_sub(generated);
5486                            if room <= 2 {
5487                                break;
5488                            }
5489                            let k = burst_k.min(room - 1);
5490                            if k < 1 {
5491                                break;
5492                            }
5493                            let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
5494                                if self
5495                                    .graph_failed
5496                                    .swap(false, std::sync::atomic::Ordering::Relaxed)
5497                                {
5498                                    self.finish_generation(&mut mtp, &mut router, true);
5499                                    return Err(
5500                                        "GPU token graph failed during greedy burst".to_string()
5501                                    );
5502                                }
5503                                break;
5504                            };
5505                            next_pos += k;
5506                            for &id in &ids {
5507                                if !commit!(id) {
5508                                    stopped = true;
5509                                    break;
5510                                }
5511                            }
5512                            if stopped {
5513                                break;
5514                            }
5515                            t_fwd = *ids.last().unwrap();
5516                        }
5517                        if stopped {
5518                            break 'decode;
5519                        }
5520                    }
5521                    // Metal: keep the draft head's cache in step through
5522                    // the trial's plain phase and a paused speculation —
5523                    // the pair (hidden, t_fwd) at next_pos−1, the step the
5524                    // round's draft 0 would take. Without it the head's
5525                    // cache lagged the trunk by every plain token for the
5526                    // rest of the generation: the batched warm-up declined
5527                    // every later round and its rows went one by one (a
5528                    // whole MTP step per accepted token), and the drafts
5529                    // attended a context with those tokens missing.
5530                    #[cfg(target_os = "macos")]
5531                    if graph_spec
5532                        && spec_watchdog_off
5533                        && next_pos > 0
5534                        && self.mtp_graph_mode == Some(true)
5535                        && crate::gpu::q1_force()
5536                    {
5537                        if let Some(m) = mtp.as_mut() {
5538                            let _ = self.mtp_step_metal(m, &hidden, t_fwd, next_pos - 1, false);
5539                        }
5540                    }
5541                    hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
5542                    next_pos += 1;
5543                    // Dynamic routing: the forward updated φ; ask the
5544                    // router whether to switch skills before the next token.
5545                    if let Some(r) = &mut router {
5546                        let phi = self.dyn_phi_ema.clone();
5547                        let decision = r.step(&phi, generated);
5548                        if let Some(new_active) = decision {
5549                            let _ = self.set_active_skill(new_active);
5550                        }
5551                        // Backfill this token's coherence + switch flag from
5552                        // the just-run eval (freshest measured values).
5553                        if trace_on {
5554                            if let Some(last) = traces.last_mut() {
5555                                let e = r.last_best_e();
5556                                last.recon = e.is_finite().then_some(e);
5557                                last.switched = decision.is_some();
5558                            }
5559                        }
5560                    }
5561                }
5562            }
5563        }
5564
5565        let cancelled = finish_reason == "cancelled";
5566        // A generation during which the router switched weights holds no
5567        // state any single overlay would produce (each switch cleared the
5568        // cache mid-sequence), so it leaves no reuse key behind either.
5569        let dyn_switched = router.as_ref().is_some_and(|r| !r.switches.is_empty());
5570        if mimo_spec {
5571            if let Some(st) = self.mimo_mtp.as_ref() {
5572                let line = st.stats.line();
5573                tracing::info!("{line}");
5574                if std::env::var_os("CMF_MIMO_MTP_STATS").is_some() {
5575                    eprintln!("{line}");
5576                }
5577            }
5578        }
5579        self.finish_generation(&mut mtp, &mut router, cancelled);
5580
5581        let output_ids = &all_ids[input_ids.len()..];
5582        // Forwarded = prompt + all generated but the LAST sampled token
5583        // (emitted without being fed back). Exact only without MTP —
5584        // reuse is gated off when MTP is active.
5585        let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
5586        // A MiMo speculative round that stopped on an accepted draft (EOS,
5587        // cancel) leaves verify rows past the committed stream in the cache:
5588        // never offer that cache for reuse. Neither does a router that
5589        // switched weights mid-sequence (no single overlay produced it).
5590        let consumed = std::mem::take(&mut all_ids);
5591        if dyn_switched {
5592            self.clear_sequence_state();
5593        } else if cancelled || mimo_spec || prompt_rows.is_some() {
5594            self.clear_history();
5595        } else {
5596            self.record_consumed_prefix(&consumed[..forwarded.min(consumed.len())], reuse_from);
5597        }
5598        all_ids = consumed;
5599        let output_ids = &all_ids[input_ids.len()..];
5600        confidence.truncate(output_ids.len()); // guard against any overshoot
5601        traces.truncate(output_ids.len());
5602        Ok(GenerateResult {
5603            text: self.tokenizer.decode(output_ids),
5604            token_ids: output_ids.to_vec(),
5605            prompt_tokens: input_ids.len(),
5606            tokens_generated: generated,
5607            finish_reason,
5608            mtp_drafted: drafted,
5609            mtp_accepted: accepted,
5610            token_confidence: confidence,
5611            traces,
5612        })
5613    }
5614
5615    /// One MTP step: feed `(hidden_p, token_{p+1})` into the draft head,
5616    /// advance its KV cache at position `p`, return the drafted token
5617    /// for position `p+2`.
5618    fn mtp_step(
5619        &mut self,
5620        m: &mut MtpModule,
5621        hidden: &[f32],
5622        next_token: u32,
5623        position: usize,
5624    ) -> u32 {
5625        self.mtp_step_h(m, hidden, next_token, position).0
5626    }
5627
5628    /// Tally for `CMF_MTP_CHAIN_PROBE`: per depth, how often the CHAIN is
5629    /// still an exact prefix of the real continuation. Printed every 128
5630    /// depth-0 samples so a killed run still shows its table.
5631    fn chain_probe_note(depth: usize, prefix_ok: bool) {
5632        use std::sync::Mutex;
5633        static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
5634        let mut t = T.lock().unwrap();
5635        if t.len() <= depth {
5636            t.resize(depth + 1, (0, 0));
5637        }
5638        t[depth].0 += 1;
5639        t[depth].1 += prefix_ok as u64;
5640        if depth == 0 && t[0].0 % 128 == 0 {
5641            let line: Vec<String> = t
5642                .iter()
5643                .enumerate()
5644                .map(|(d, (n, k))| {
5645                    format!(
5646                        "d{}={:.0}%({n})",
5647                        d + 1,
5648                        100.0 * *k as f64 / (*n).max(1) as f64
5649                    )
5650                })
5651                .collect();
5652            eprintln!("mtp-chain: {}", line.join(" "));
5653        }
5654    }
5655
5656    /// `mtp_step` that also hands back the block's own output hidden — the
5657    /// state a CHAINED draft feeds the next step, the way a multi-token
5658    /// speculative round iterates the head on itself.
5659    /// One MTP block step from (trunk hidden, token): the head's LOGITS
5660    /// and the block's own hidden for chaining. The draft is argmax of the
5661    /// logits on the greedy path and a draw from their post-chain
5662    /// distribution on the sampling path.
5663    fn mtp_step_hl(
5664        &mut self,
5665        m: &mut MtpModule,
5666        hidden: &[f32],
5667        next_token: u32,
5668        position: usize,
5669    ) -> (Vec<f32>, Vec<f32>) {
5670        // The graph arm: the MTP block as a one-layer token graph with the
5671        // head fused — device attention over the block's own KV mirror,
5672        // one submit for block + head, hidden and logits back together.
5673        // Decided once per generation (see `mtp_graph_mode`).
5674        #[cfg(target_os = "macos")]
5675        if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
5676            if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
5677                self.mtp_graph_mode = Some(true);
5678                return r;
5679            }
5680            if self.mtp_graph_mode == Some(true) {
5681                tracing::error!("mtp Metal graph failed after admission");
5682                self.clear_sequence_state();
5683                self.graph_failed
5684                    .store(true, std::sync::atomic::Ordering::Relaxed);
5685                self.cancel
5686                    .store(true, std::sync::atomic::Ordering::Relaxed);
5687                return (Vec::new(), Vec::new());
5688            }
5689            self.mtp_graph_mode = Some(false);
5690        }
5691        #[cfg(feature = "gpu")]
5692        if self.mtp_graph_mode != Some(false) {
5693            if !self.mtp_graph_ok(m) {
5694                if self.mtp_graph_mode == Some(true) {
5695                    // A mirror was already admitted, so a capability change
5696                    // cannot safely switch this request to the stale CPU
5697                    // cache.  Keep the same terminal contract as a failed
5698                    // token graph.
5699                    tracing::error!("mtp graph became unavailable after admission");
5700                    self.clear_sequence_state();
5701                    self.graph_failed
5702                        .store(true, std::sync::atomic::Ordering::Relaxed);
5703                    self.cancel
5704                        .store(true, std::sync::atomic::Ordering::Relaxed);
5705                    return (Vec::new(), Vec::new());
5706                }
5707                self.mtp_graph_mode = Some(false);
5708            } else {
5709                if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
5710                    self.mtp_graph_mode = Some(true);
5711                    return r;
5712                }
5713                if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
5714                    // A token graph can have admitted a persistent MTP/GDN
5715                    // mirror before its readback failed.  The CPU MTP cache
5716                    // is not a valid continuation in that state; leave the
5717                    // flag set so the generation caller returns through its
5718                    // terminal error path instead of silently switching
5719                    // arithmetic.
5720                    return (Vec::new(), Vec::new());
5721                }
5722                // `mtp_graph_ok` was true, so a None here means a refusal or
5723                // failure after graph admission.  Do not fall through to a
5724                // CPU cache whose rows may lag the device mirror.
5725                tracing::error!("mtp graph failed or declined after admission");
5726                self.clear_sequence_state();
5727                self.graph_failed
5728                    .store(true, std::sync::atomic::Ordering::Relaxed);
5729                self.cancel
5730                    .store(true, std::sync::atomic::Ordering::Relaxed);
5731                return (Vec::new(), Vec::new());
5732            }
5733        }
5734        // fc concat order is [enorm(embed); hnorm(hidden)] — EMBEDDING
5735        // FIRST. Verified by the oracle (converter/mtp_oracle.py):
5736        // [emb;hid] → 45.8% acceptance, [hid;emb] → 0.00%.
5737        let e = self.embed_single(next_token);
5738        let mut cat = vec![0.0f32; 2 * self.hidden_size];
5739        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5740        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5741        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5742        let mut x = vec![0.0f32; self.hidden_size];
5743        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5744
5745        // One standard transformer block over the MTP's own cache.
5746        let lw = &m.layer;
5747        inference::rms_norm_into(
5748            &x,
5749            &lw.input_norm,
5750            self.rms_eps,
5751            self.norm_style,
5752            &mut self.ws.n1,
5753        );
5754        let attn = match &lw.attn {
5755            // MLA models carry no MTP head; this path cannot see them.
5756            AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
5757            AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
5758            AttnKind::Full {
5759                wq,
5760                wk,
5761                wv,
5762                wo,
5763                q_norm,
5764                k_norm,
5765                output_gate,
5766                softplus_gate,
5767                bias,
5768            } => {
5769                let mut cfg = self.attn_cfg(position);
5770                cfg.q_norm = q_norm.as_deref();
5771                cfg.k_norm = k_norm.as_deref();
5772                cfg.output_gate = *output_gate;
5773                cfg.softplus_gate = softplus_gate
5774                    .as_ref()
5775                    .map(|(gate, per_head)| (gate, *per_head));
5776                cfg.bias = bias
5777                    .as_ref()
5778                    .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
5779                attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
5780            }
5781            AttnKind::Linear(_)
5782            | AttnKind::LinearGdn(_)
5783            | AttnKind::ShortConv(_)
5784            | AttnKind::Bounded(_) => {
5785                unreachable!("MTP block is full attention")
5786            }
5787        };
5788        for (i, &a) in attn.iter().enumerate() {
5789            x[i] += a;
5790        }
5791        inference::rms_norm_into(
5792            &x,
5793            &lw.post_norm,
5794            self.rms_eps,
5795            self.norm_style,
5796            &mut self.ws.p1,
5797        );
5798        let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
5799        for (i, &f) in ffn.iter().enumerate() {
5800            x[i] += f;
5801        }
5802
5803        inference::rms_norm_into(
5804            &x,
5805            &m.final_norm,
5806            self.rms_eps,
5807            self.norm_style,
5808            &mut self.ws.n1,
5809        );
5810        let lg = self.lm_head_forward(&self.ws.n1);
5811        (lg, x)
5812    }
5813
5814    /// `mtp_step_hl` reduced to the greedy draft: argmax of the head.
5815    fn mtp_step_h(
5816        &mut self,
5817        m: &mut MtpModule,
5818        hidden: &[f32],
5819        next_token: u32,
5820        position: usize,
5821    ) -> (u32, Vec<f32>) {
5822        let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
5823        let draft = sampler::argmax(&lg);
5824        attention::recycle_buf(&mut lg);
5825        (draft, x)
5826    }
5827
5828    /// One speculative round for the trial: rounds 1..5 of a `Spec` phase
5829    /// advance it (the monitor already averaged this round); after five,
5830    /// the plain phase runs (once — a known plain rate decides at once);
5831    /// a decided speculation keeps re-checking the rule every round and
5832    /// stops after four losing rounds in a row.
5833    fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
5834        match trial {
5835            SpecTrial::Spec { t0, gen0, rounds } => {
5836                let rounds = rounds + 1;
5837                if rounds >= 5 {
5838                    if mon.plain_ms > 0.0 {
5839                        let keep = mon.pays();
5840                        mon.fails = 0;
5841                        tracing::info!(
5842                            "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5843                            mon.tokens,
5844                            mon.round_ms,
5845                            mon.plain_ms,
5846                            if keep { "speculating" } else { "plain" }
5847                        );
5848                        SpecTrial::Decided {
5849                            spec: keep,
5850                            recheck_at: if keep { usize::MAX } else { generated + 128 },
5851                        }
5852                    } else if mon.pays() {
5853                        // Metal: the rounds land enough tokens each that no
5854                        // plain measurement is needed — keep speculating,
5855                        // and re-check every round (a losing streak sends
5856                        // the loop to the plain phase, below).
5857                        mon.fails = 0;
5858                        tracing::info!(
5859                            "speculation trial: {:.2} tok/round in {:.1} ms — speculating (plain not timed)",
5860                            mon.tokens,
5861                            mon.round_ms,
5862                        );
5863                        SpecTrial::Decided {
5864                            spec: true,
5865                            recheck_at: usize::MAX,
5866                        }
5867                    } else {
5868                        SpecTrial::Plain {
5869                            t0: std::time::Instant::now(),
5870                            gen0: generated,
5871                        }
5872                    }
5873                } else {
5874                    SpecTrial::Spec { t0, gen0, rounds }
5875                }
5876            }
5877            SpecTrial::Decided { spec: true, .. } => {
5878                if mon.pays() {
5879                    mon.fails = 0;
5880                    trial
5881                } else {
5882                    mon.fails += 1;
5883                    if mon.fails >= 4 {
5884                        if mon.plain_ms <= 0.0 {
5885                            // Metal, plain never timed: four doubtful rounds
5886                            // buy the (bounded) plain measurement, and the
5887                            // exact rule decides from it.
5888                            tracing::info!(
5889                                "speculation doubtful: {:.2} tok/round in {:.1} ms — timing plain",
5890                                mon.tokens,
5891                                mon.round_ms,
5892                            );
5893                            return SpecTrial::Plain {
5894                                t0: std::time::Instant::now(),
5895                                gen0: generated,
5896                            };
5897                        }
5898                        tracing::info!(
5899                            "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
5900                            mon.tokens,
5901                            mon.round_ms,
5902                            mon.plain_ms
5903                        );
5904                        SpecTrial::Decided {
5905                            spec: false,
5906                            recheck_at: generated + 128,
5907                        }
5908                    } else {
5909                        trial
5910                    }
5911                }
5912            }
5913            other => other,
5914        }
5915    }
5916
5917    /// The MTP block's device-mirror id: the trunk's id with a high bit,
5918    /// so the (kv_id, layer) mirror keys never collide.
5919    fn mtp_kv_id(&self) -> u64 {
5920        self.graph_kv_id | (1u64 << 40)
5921    }
5922
5923    /// The MTP block's mirror layer index: 0 — its own kv_id keeps it
5924    /// apart from the trunk, and the BATCH graph (the warm-up path) keys
5925    /// its mirrors at layer 0 with no base of its own, so the draft's
5926    /// token graph must key the same slot.
5927    const MTP_LAYER_BASE: usize = 0;
5928
5929    /// The wgpu MTP draft writes speculative rows straight into its device
5930    /// mirror while the CPU owner retains only the real prompt/decode anchor.
5931    /// After verification, move that mirror cursor back to the anchor before
5932    /// replaying accepted pairs.  The next graph append then sees the same
5933    /// contiguous position as the CPU/Metal path without uploading stale
5934    /// speculative rows.
5935    #[cfg(feature = "gpu")]
5936    fn rewind_mtp_graph_mirror(&self, stored: usize) -> bool {
5937        self.mtp_graph_mode != Some(true)
5938            || crate::gpu::graph_kv_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, stored)
5939    }
5940
5941    /// A speculative verify graph appends the full `k+1` trunk rows before
5942    /// the acceptance count is known.  GDN state already has a snapshot
5943    /// restore; Full-attention mirrors need the matching logical cursor
5944    /// rewind so the next graph call does not reject an ahead-of-position KV
5945    /// cache after a partial acceptance.
5946    #[cfg(feature = "gpu")]
5947    fn rewind_trunk_graph_mirrors(&self, stored: usize) -> bool {
5948        let mut ok = true;
5949        let mut expected = false;
5950        for li in 0..self.num_layers {
5951            if matches!(
5952                self.weights.layers[self.phys_layer(li)].attn,
5953                AttnKind::Full { .. }
5954            ) {
5955                expected = true;
5956                ok &= crate::gpu::graph_kv_set_stored(self.graph_kv_id, li, stored);
5957            }
5958        }
5959        !expected || ok
5960    }
5961
5962    /// Count the recurrent layers participating in the trunk verify graph.
5963    /// Snapshot restore is all-or-nothing across that set; deriving the count
5964    /// from the model keeps the restore contract valid for looped models too.
5965    fn graph_gdn_layer_count(&self) -> usize {
5966        (0..self.num_layers)
5967            .filter(|&li| {
5968                matches!(
5969                    &self.weights.layers[self.phys_layer(li)].attn,
5970                    AttnKind::LinearGdn(_)
5971                )
5972            })
5973            .count()
5974    }
5975
5976    /// The block's input from (trunk hidden, token): eh_proj · [enorm(e);
5977    /// hnorm(h)] — the same arithmetic the per-op path starts with.
5978    fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
5979        let e = self.embed_single(next_token);
5980        let mut cat = vec![0.0f32; 2 * self.hidden_size];
5981        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5982        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5983        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5984        let mut x = vec![0.0f32; self.hidden_size];
5985        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5986        x
5987    }
5988
5989    /// Is the MTP block graphable at all (device up, full attention
5990    /// without softplus, dense FFN)? The plan itself is built per call.
5991    #[cfg(feature = "gpu")]
5992    fn mtp_block_graph_ok(&self, m: &MtpModule) -> bool {
5993        if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
5994            return false;
5995        }
5996        if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
5997            || !crate::gpu::enabled_here()
5998            || self.attn_softcap > 0.0
5999            || self.attention_heads_per_layer.is_some()
6000            // The block graph caches V as wide as K and feeds o_proj
6001            // nh·head_dim; a narrow-V model keeps its MTP block per-op.
6002            || self.v_head_dim.is_some()
6003        {
6004            return false;
6005        }
6006        matches!(
6007            &m.layer.attn,
6008            AttnKind::Full {
6009                softplus_gate: None,
6010                ..
6011            }
6012        ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
6013    }
6014
6015    /// Full MTP token-graph eligibility, including the fused lm-head and all
6016    /// block projection weights.  Keep this distinct from the block-only
6017    /// check: prompt warm-up does not need the head, while a draft step does.
6018    #[cfg(feature = "gpu")]
6019    fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
6020        if !self.mtp_block_graph_ok(m) {
6021            return false;
6022        }
6023        let AttnKind::Full { wq, wk, wv, wo, .. } = &m.layer.attn else {
6024            return false;
6025        };
6026        let FfnKind::Dense(d) = &m.layer.ffn else {
6027            return false;
6028        };
6029        d.segs.is_empty()
6030            && wq.graph_weight().is_some()
6031            && wk.graph_weight().is_some()
6032            && wv.graph_weight().is_some()
6033            && wo.graph_weight().is_some()
6034            && d.gate_proj.graph_weight().is_some()
6035            && d.up_proj.graph_weight().is_some()
6036            && d.down_proj.graph_weight().is_some()
6037            && self.weights.lm_head.graph_weight().is_some()
6038    }
6039
6040    /// One MTP block step on the wgpu token graph: block + fused head in
6041    /// one submit, the block hidden and the logits read back together.
6042    /// None = the graph cannot take this block (softplus gate, non-dense
6043    /// FFN, unquantized head, no device) — the caller keeps the per-op
6044    /// path for the whole generation.
6045    #[cfg(feature = "gpu")]
6046    fn mtp_step_graph(
6047        &mut self,
6048        m: &mut MtpModule,
6049        hidden: &[f32],
6050        next_token: u32,
6051        position: usize,
6052    ) -> Option<(Vec<f32>, Vec<f32>)> {
6053        if !self.mtp_graph_ok(m) {
6054            return None;
6055        }
6056        let lw = &m.layer;
6057        let AttnKind::Full {
6058            wq,
6059            wk,
6060            wv,
6061            wo,
6062            q_norm,
6063            k_norm,
6064            output_gate,
6065            softplus_gate,
6066            bias,
6067        } = &lw.attn
6068        else {
6069            return None;
6070        };
6071        if softplus_gate.is_some() {
6072            return None;
6073        }
6074        let FfnKind::Dense(d) = &lw.ffn else {
6075            return None;
6076        };
6077        if !d.segs.is_empty() {
6078            return None; // tube layers run on the segmented path
6079        }
6080        // The block's input first: it borrows `self` mutably (embed scratch,
6081        // pool), the plan below borrows the weights immutably.
6082        let mut x = self.mtp_block_input(m, hidden, next_token);
6083        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6084            let (_, i, kind, rs) = t.graph_weight()?;
6085            Some(crate::gpu::GraphW {
6086                idx: i,
6087                kind,
6088                row_scale: rs,
6089                data: &[],
6090                prism: crate::gpu::GraphPrismOp::None,
6091                affine: false,
6092            })
6093        }
6094        let (model, _, _, _) = wq.graph_weight()?;
6095        let model = model.clone();
6096        let (lm_gw, lm_rows) = {
6097            let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6098            // The draft's head over the CMF_DRAFT_VOCAB shortlist (the same
6099            // cut the native Metal draft takes): 662 MB a step on Qwen3.8
6100            // becomes 170 MB at 65536; the verify keeps the full head.
6101            let rows = if kind == 6 {
6102                self.draft_head_rows(self.weights.lm_head.rows())
6103            } else {
6104                self.weights.lm_head.rows()
6105            };
6106            (
6107                crate::gpu::GraphW {
6108                    idx: i,
6109                    kind,
6110                    row_scale: rs,
6111                    data: &[],
6112                    prism: crate::gpu::GraphPrismOp::None,
6113                    affine: false,
6114                },
6115                rows,
6116            )
6117        };
6118        let layer = crate::gpu::GraphLayer {
6119            input_norm: &lw.input_norm,
6120            attn: crate::gpu::GraphAttn::Full {
6121                wq: gw(wq)?,
6122                wk: gw(wk)?,
6123                wv: gw(wv)?,
6124                wo: gw(wo)?,
6125                q_norm: q_norm.as_deref(),
6126                k_norm: k_norm.as_deref(),
6127                late_qk_norm: self.qk_norm_after_rope,
6128                bias: bias
6129                    .as_ref()
6130                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6131                output_gate: *output_gate,
6132                cpu_k: m.kv.k_heads(),
6133                cpu_v: m.kv.v_heads(),
6134                geom: None,
6135            },
6136            post_norm: &lw.post_norm,
6137            ffn: crate::gpu::GraphFfn::Dense {
6138                gate: gw(&d.gate_proj)?,
6139                up: gw(&d.up_proj)?,
6140                down: gw(&d.down_proj)?,
6141            },
6142        };
6143        let nh = self.num_heads;
6144        let (nkv, hd, rd) = self.layer_geom(0);
6145        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6146        let mut logits = Vec::new();
6147        let ok = crate::gpu::forward_token_graph(
6148            &model,
6149            self.mtp_kv_id(),
6150            std::slice::from_ref(&layer),
6151            &[None],
6152            self.o1_epoch,
6153            &self.inv_freq,
6154            &mut x,
6155            nh,
6156            nkv,
6157            hd,
6158            self.attn_scale,
6159            rd,
6160            self.hidden_size,
6161            self.intermediate_size,
6162            position,
6163            self.kv_cache.max_seq_len,
6164            gemma,
6165            self.rms_eps as f32,
6166            Some((&lm_gw, lm_rows)),
6167            &m.final_norm,
6168            &mut logits,
6169            &[],
6170            1,
6171            None,
6172            None,
6173            None,
6174            Self::MTP_LAYER_BASE,
6175            true,
6176        );
6177        match ok {
6178            crate::gpu::TokenGraphOutcome::Completed => {}
6179            crate::gpu::TokenGraphOutcome::Declined => return None,
6180            crate::gpu::TokenGraphOutcome::Failed => {
6181                // The backend has already admitted persistent state.  Keep
6182                // this distinct from a capability refusal so the caller
6183                // cannot switch to the stale CPU MTP cache.
6184                self.clear_sequence_state();
6185                self.graph_failed
6186                    .store(true, std::sync::atomic::Ordering::Relaxed);
6187                self.cancel
6188                    .store(true, std::sync::atomic::Ordering::Relaxed);
6189                return None;
6190            }
6191        }
6192        logits.resize(self.vocab_size, 0.0);
6193        Some((logits, x))
6194    }
6195
6196    /// The warm-ups of one speculative round on the device: every accepted
6197    /// (hidden, token) pair as ONE batched graph run over the MTP block
6198    /// (no head) — its kv_append lands the pairs in the block's mirror.
6199    /// `pairs` are consecutive positions from `first_pos`.  The tri-state
6200    /// result is intentional: a refusal before admission may use the
6201    /// per-row/CPU route, while a failure after admission must terminate the
6202    /// sequence rather than fall through to a stale CPU cache.
6203    #[cfg(feature = "gpu")]
6204    fn mtp_warm_graph(
6205        &mut self,
6206        m: &mut MtpModule,
6207        pairs: &[(&[f32], u32)],
6208        first_pos: usize,
6209    ) -> crate::gpu::BatchGraphOutcome {
6210        if pairs.is_empty() {
6211            return crate::gpu::BatchGraphOutcome::Completed;
6212        }
6213        if !self.mtp_block_graph_ok(m) {
6214            return crate::gpu::BatchGraphOutcome::Declined;
6215        }
6216        let hs = self.hidden_size;
6217        // Block inputs for every pair (eh_proj on the per-op path, one
6218        // matvec each — the plan's own prologue).
6219        let mut hiddens = Vec::with_capacity(pairs.len() * hs);
6220        for (h, t) in pairs {
6221            hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
6222        }
6223        let lw = &m.layer;
6224        let AttnKind::Full {
6225            wq,
6226            wk,
6227            wv,
6228            wo,
6229            q_norm,
6230            k_norm,
6231            output_gate,
6232            bias,
6233            ..
6234        } = &lw.attn
6235        else {
6236            return crate::gpu::BatchGraphOutcome::Declined;
6237        };
6238        let FfnKind::Dense(d) = &lw.ffn else {
6239            return crate::gpu::BatchGraphOutcome::Declined;
6240        };
6241        if !d.segs.is_empty() {
6242            return crate::gpu::BatchGraphOutcome::Declined; // tube layers run on the segmented path
6243        }
6244        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6245            let (_, i, kind, rs) = t.graph_weight()?;
6246            Some(crate::gpu::GraphW {
6247                idx: i,
6248                kind,
6249                row_scale: rs,
6250                data: &[],
6251                prism: crate::gpu::GraphPrismOp::None,
6252                affine: false,
6253            })
6254        }
6255        let Some((model, _, _, _)) = wq.graph_weight() else {
6256            return crate::gpu::BatchGraphOutcome::Declined;
6257        };
6258        let model = model.clone();
6259        let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
6260            gw(wq),
6261            gw(wk),
6262            gw(wv),
6263            gw(wo),
6264            gw(&d.gate_proj),
6265            gw(&d.up_proj),
6266            gw(&d.down_proj),
6267        ) else {
6268            return crate::gpu::BatchGraphOutcome::Declined;
6269        };
6270        let layer = crate::gpu::GraphLayer {
6271            input_norm: &lw.input_norm,
6272            attn: crate::gpu::GraphAttn::Full {
6273                wq: gwq,
6274                wk: gwk,
6275                wv: gwv,
6276                wo: gwo,
6277                q_norm: q_norm.as_deref(),
6278                k_norm: k_norm.as_deref(),
6279                late_qk_norm: self.qk_norm_after_rope,
6280                bias: bias
6281                    .as_ref()
6282                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6283                output_gate: *output_gate,
6284                cpu_k: m.kv.k_heads(),
6285                cpu_v: m.kv.v_heads(),
6286                geom: None,
6287            },
6288            post_norm: &lw.post_norm,
6289            ffn: crate::gpu::GraphFfn::Dense {
6290                gate: gg,
6291                up: gu,
6292                down: gd,
6293            },
6294        };
6295        let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
6296        let nh = self.num_heads;
6297        let (nkv, hd, rd) = self.layer_geom(0);
6298        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6299        crate::gpu::forward_batch_graph(
6300            &model,
6301            self.mtp_kv_id(),
6302            std::slice::from_ref(&layer),
6303            &self.inv_freq,
6304            &mut hiddens,
6305            nh,
6306            nkv,
6307            hd,
6308            rd,
6309            hs,
6310            self.intermediate_size,
6311            &positions,
6312            self.kv_cache.max_seq_len,
6313            gemma,
6314            self.rms_eps as f32,
6315            self.attn_scale,
6316            pairs.len(),
6317            &[],
6318            0,
6319            None,
6320            None,
6321        )
6322    }
6323
6324    /// Complete an MTP warm-up after the batched graph has refused.  A
6325    /// graphable block is retried one row at a time; once any device row has
6326    /// been admitted, a CPU fallback would observe a stale mirror, so every
6327    /// token-graph refusal is terminal.  If the block is not graphable and no
6328    /// mirror exists yet, warming on the CPU is safe and records the CPU mode
6329    /// for the rest of the generation.
6330    #[cfg(feature = "gpu")]
6331    fn mtp_warm_graph_fallback(
6332        &mut self,
6333        m: &mut MtpModule,
6334        pairs: &[(&[f32], u32)],
6335        first_pos: usize,
6336    ) -> bool {
6337        if pairs.is_empty() {
6338            return true;
6339        }
6340        let graphable = self.mtp_block_graph_ok(m);
6341        if !graphable {
6342            // A previously admitted mirror cannot be made coherent by
6343            // appending to the host cache.  The caller turns this into a
6344            // terminal generation error and clears both mirrors.
6345            if self.mtp_graph_mode == Some(true) {
6346                return false;
6347            }
6348            self.mtp_graph_mode = Some(false);
6349            for (j, (h, t)) in pairs.iter().enumerate() {
6350                self.mtp_warm(m, h, *t, first_pos + j);
6351            }
6352            return true;
6353        }
6354
6355        // The batch refusal is recoverable only through the same device
6356        // state.  Keep rows owned until each token graph has completed; a
6357        // None is treated as unsafe because the token-graph API deliberately
6358        // collapses its backend refusal/failure into that result.
6359        for (j, (h, t)) in pairs.iter().enumerate() {
6360            if self.mtp_step_graph(m, h, *t, first_pos + j).is_none() {
6361                return false;
6362            }
6363        }
6364        self.mtp_graph_mode = Some(true);
6365        true
6366    }
6367
6368    /// Warm a contiguous set of MTP pairs using the existing graph seam, with
6369    /// an all-or-nothing error contract for callers that already admitted the
6370    /// trunk batch.  The non-GPU build keeps the same pair accounting while
6371    /// using the established CPU warm path.
6372    #[cfg(feature = "gpu")]
6373    fn mtp_warm_prefill_pairs(
6374        &mut self,
6375        m: &mut MtpModule,
6376        pairs: &[(&[f32], u32)],
6377        first_pos: usize,
6378    ) -> Result<(), &'static str> {
6379        // Keep unsupported token-graph heads on the established CPU MTP
6380        // route before admitting any block mirror.  Once a device mirror is
6381        // active, the same condition is terminal because CPU rows cannot
6382        // repair its state.
6383        if self.mtp_graph_mode == Some(false) || !self.mtp_graph_ok(m) {
6384            if self.mtp_graph_mode == Some(true) {
6385                return Err("MTP token graph became unavailable after admission");
6386            }
6387            self.mtp_graph_mode = Some(false);
6388            for (j, (h, t)) in pairs.iter().enumerate() {
6389                self.mtp_warm(m, h, *t, first_pos + j);
6390            }
6391            return Ok(());
6392        }
6393        match self.mtp_warm_graph(m, pairs, first_pos) {
6394            crate::gpu::BatchGraphOutcome::Completed => {
6395                if !pairs.is_empty() {
6396                    self.mtp_graph_mode = Some(true);
6397                }
6398                Ok(())
6399            }
6400            crate::gpu::BatchGraphOutcome::Declined => {
6401                if self.mtp_warm_graph_fallback(m, pairs, first_pos) {
6402                    Ok(())
6403                } else {
6404                    Err("MTP warm-up fallback failed after device admission")
6405                }
6406            }
6407            crate::gpu::BatchGraphOutcome::Failed => {
6408                Err("MTP warm batch graph failed after admission")
6409            }
6410        }
6411    }
6412
6413    #[cfg(not(feature = "gpu"))]
6414    fn mtp_warm_prefill_pairs(
6415        &mut self,
6416        m: &mut MtpModule,
6417        pairs: &[(&[f32], u32)],
6418        first_pos: usize,
6419    ) -> Result<(), &'static str> {
6420        for (j, (h, t)) in pairs.iter().enumerate() {
6421            self.mtp_warm(m, h, *t, first_pos + j);
6422        }
6423        Ok(())
6424    }
6425
6426    /// The MTP block alone — advance its KV with a (hidden, token) pair the
6427    /// verify just proved, without paying the head. What keeps the draft's
6428    /// attention context warm between speculative rounds.
6429    fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
6430        let e = self.embed_single(next_token);
6431        let mut cat = vec![0.0f32; 2 * self.hidden_size];
6432        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6433        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6434        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6435        let mut x = vec![0.0f32; self.hidden_size];
6436        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6437        inference::rms_norm_into(
6438            &x,
6439            &m.layer.input_norm,
6440            self.rms_eps,
6441            self.norm_style,
6442            &mut self.ws.n1,
6443        );
6444        let attn = match &m.layer.attn {
6445            AttnKind::Full {
6446                wq,
6447                wk,
6448                wv,
6449                wo,
6450                q_norm,
6451                k_norm,
6452                output_gate,
6453                softplus_gate,
6454                bias,
6455            } => {
6456                let mut cfg = self.attn_cfg(position);
6457                cfg.q_norm = q_norm.as_deref();
6458                cfg.k_norm = k_norm.as_deref();
6459                cfg.output_gate = *output_gate;
6460                cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
6461                cfg.bias = bias
6462                    .as_ref()
6463                    .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
6464                attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
6465            }
6466            _ => return,
6467        };
6468        let _ = attn;
6469    }
6470
6471    /// Speculative decode ON the wgpu whole-token graph: draft k with the
6472    /// MTP head, verify all of them plus the tip in ONE batched graph
6473    /// submit whose tail folds the head, commit the accepted prefix and
6474    /// roll the GDN state back to the last real position. Greedy only —
6475    /// output equals the plain graph's token for token, the way the DSV4
6476    /// verify equals the walk.
6477    #[cfg(feature = "gpu")]
6478    #[allow(clippy::too_many_arguments)]
6479    fn graph_spec_step(
6480        &mut self,
6481        m: &mut MtpModule,
6482        hidden: &[f32],
6483        t_next: u32,
6484        next_pos: usize,
6485        drafted: &mut usize,
6486        accepted: &mut usize,
6487        // The committed stream (prompt + generated so far, `t_next`
6488        // included): the sampler chain's penalties read it, and the
6489        // sampling arm extends it with the drafts position by position.
6490        all_ids: &mut Vec<u32>,
6491        // Tokens left before `max_tokens`. A round commits up to k
6492        // accepted drafts, and those positions are already in the cache,
6493        // so the depth is capped here — trimming the output afterwards
6494        // would leave cache rows the committed stream does not have.
6495        room: usize,
6496    ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
6497        // 3 is the measured optimum on Qwen3.6-27B / RTX 5090 (medians
6498        // of three, greedy): 51.1 tok/s against a plain 49.4, where k=2
6499        // gives 46.1, k=4 50.0, k=5 47.4, k=6 45.2. Acceptance is 89-91%
6500        // throughout — what turns the curve over is the verify, which
6501        // costs ~7.4 ms per extra position, and the draft ~3 ms a step.
6502        // 4 since the draft moved onto the graph (Qwen3.8-27B / 5090:
6503        // k=3 51.2, k=4 51.8 with the per-op draft; the graph draft
6504        // halves the draft cost, so the extra draft is cheaper still).
6505        // 5 with the int8 verify (the default: measured 76.5 against
6506        // k=4's 72-74 and k=6's 74 on the 5090), 4 with the f32 one.
6507        #[cfg(target_os = "macos")]
6508        let metal_native = crate::gpu::q1_force();
6509        #[cfg(not(target_os = "macos"))]
6510        let metal_native = false;
6511        #[cfg(feature = "gpu")]
6512        let k_default = if metal_native {
6513            // the Metal verify's GEMM tile is 8 rows wide and flat in b:
6514            // seven drafts + the tip fill it for free
6515            7
6516        } else if crate::gpu_wgpu::verify_i8_on() {
6517            5
6518        } else {
6519            4
6520        };
6521        #[cfg(not(feature = "gpu"))]
6522        let k_default = 4;
6523        let k_env: Option<usize> = std::env::var("CMF_GRAPH_SPEC_K")
6524            .ok()
6525            .and_then(|v| v.parse().ok())
6526            .filter(|&v| (1..=8).contains(&v));
6527        // Adaptive depth: start below the card's flat-verify optimum and
6528        // let the accepted fraction move it — predictable text climbs to
6529        // the old default within a few rounds, prose settles at 2-3 where
6530        // the shorter verify pays.
6531        let (k_start, k_max) = if metal_native { (7, 7) } else { (3, k_default.max(5)) };
6532        let k_full: usize = k_env.unwrap_or_else(|| self.spec_k_adapt.unwrap_or(k_start));
6533        let k_spec = k_full.min(room).max(1);
6534        // a tail round cut short by `room` says nothing about the text:
6535        // it must not move the adaptive depth the next request starts at
6536        let k_capped = k_spec < k_full;
6537        if next_pos == 0 {
6538            return None;
6539        }
6540        let t_round = std::time::Instant::now();
6541        // Submissions per phase — and they say where the round's money is.
6542        // Qwen3.6-27B on an RTX 5090, k=3:
6543        //
6544        //   draft   9.3 ms / 12 submissions   (four per MTP step)
6545        //   verify 52.8 ms /  1               (the batched graph)
6546        //   commit  5.4 ms /  6               (two per warm)
6547        //
6548        // The verify is already one submit. The draft's own work is 834 MB
6549        // a step — 0.8 ms at this card's measured 1056 GB/s — against 3.1
6550        // ms measured, so ~0.58 ms of every step is round trip, not
6551        // arithmetic, and the same holds for the warms. Eighteen round
6552        // trips a round at roughly half a millisecond each is ~11 ms of a
6553        // 68 ms round: fusing the MTP block into ONE submit the way the
6554        // trunk already is projects to ~64 tok/s against today's 50.9.
6555        // That is the largest measured item left on this path.
6556        let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
6557        let sub0 = subs();
6558        // Greedy without penalties verifies by argmax equality (bit-exact
6559        // against the plain path). Anything else is speculative SAMPLING:
6560        // each draft is a DRAW from the MTP head's post-chain distribution
6561        // q_j, kept for the accept test; the verify's rows give p_j.
6562        let cfg = self.sampler_config.clone();
6563        let penalized = !(cfg.repetition_penalty == 1.0
6564            && cfg.presence_penalty == 0.0
6565            && cfg.suppress_tokens.is_empty());
6566        // Three verify regimes: plain greedy (argmax of the raw rows),
6567        // greedy WITH penalties (argmax of the penalized rows — a single
6568        // pass each, no distributions), and sampling (draw / accept /
6569        // correct on post-chain distributions).
6570        let greedy_pen = cfg.temperature < 1e-6 && penalized;
6571        let sampling = cfg.temperature >= 1e-6;
6572        // Sampling with a top-k goes through the SPARSE chain: the dense
6573        // one builds nine 248k-float distributions a round (four drafts,
6574        // five verify rows) and measured 19-22 tok/s against a plain 40 —
6575        // the host, not the card. Sparse, the same nine cost tens of
6576        // microseconds each.
6577        let sparse = sampling && sampler::sparse_ok(&cfg);
6578        let base_len = all_ids.len();
6579        if sampling && !sparse && self.spec_q.len() < k_spec {
6580            self.spec_q.resize_with(k_spec, Vec::new);
6581        }
6582        if sparse && self.spec_qs.len() < k_spec {
6583            self.spec_qs.resize_with(k_spec, Vec::new);
6584        }
6585        // Draft the chain: first from the trunk's tip hidden, then the head
6586        // iterating on itself. Rows land in the MTP KV; the chain rows past
6587        // the first are speculation over speculative state and roll back
6588        // below, replaced by verified pairs.
6589        let mut drafts = Vec::with_capacity(k_spec);
6590        let mut hx = hidden.to_vec();
6591        // CMF_SPEC_DBG=1: draft 0 through BOTH MTP arms (graph and per-op)
6592        // from the same inputs — are the arms the difference, or the inputs?
6593        let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
6594        spec_stamp("pro");
6595        // Plain greedy on native Metal: the whole chain as one command
6596        // buffer (device argmax + embedding gather between the steps).
6597        // A decline before commit hands the round to the per-step loop
6598        // below; a failure after commit is terminal, like any graph
6599        // failure after admission.
6600        #[cfg(target_os = "macos")]
6601        if metal_native && !sampling && !greedy_pen && self.mtp_graph_mode != Some(false) {
6602            match self.mtp_draft_chain_metal(m, hidden, t_next, next_pos - 1, k_spec) {
6603                Ok(ids) => {
6604                    self.mtp_graph_mode = Some(true);
6605                    drafts = ids;
6606                }
6607                Err(true) => {
6608                    tracing::error!("mtp Metal draft chain failed after commit");
6609                    self.clear_sequence_state();
6610                    self.graph_failed
6611                        .store(true, std::sync::atomic::Ordering::Relaxed);
6612                    self.cancel
6613                        .store(true, std::sync::atomic::Ordering::Relaxed);
6614                    return None;
6615                }
6616                Err(false) => {}
6617            }
6618        }
6619        for j in drafts.len()..k_spec {
6620            let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
6621            let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
6622            if spec_dbg {
6623                let saved = self.mtp_graph_mode;
6624                self.mtp_graph_mode = Some(false);
6625                let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6626                self.mtp_graph_mode = saved;
6627                if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6628                    return None;
6629                }
6630                m.kv.truncate_last(1);
6631                dbg_ref = Some(r);
6632            }
6633            let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6634            if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6635                return None;
6636            }
6637            if let Some((lg_cpu, h_cpu)) = dbg_ref {
6638                let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
6639                let dl = lg
6640                    .iter()
6641                    .zip(&lg_cpu)
6642                    .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6643                let dh = hj
6644                    .iter()
6645                    .zip(&h_cpu)
6646                    .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6647                eprintln!(
6648                    "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 {}",
6649                    next_pos - 1 + j,
6650                    sampler::argmax(&lg_cpu),
6651                    sampler::argmax(&lg),
6652                    n(&h_cpu),
6653                    n(&hj),
6654                    m.kv.seq_len
6655                );
6656            }
6657            let dj = if sparse {
6658                let mut q = std::mem::take(&mut self.spec_qs[j]);
6659                let ok = sampler::sparse_distribution_into(
6660                    &lg,
6661                    &cfg,
6662                    all_ids,
6663                    &mut self.sampler_scratch,
6664                    self.pool.as_deref(),
6665                    &mut q,
6666                );
6667                let d = if ok {
6668                    sampler::draw_sparse(&q, &mut self.rng)
6669                } else {
6670                    // everything filtered: the dense chain's greedy fallback
6671                    let t = sampler::argmax(&lg);
6672                    q.clear();
6673                    q.push((t, 1.0));
6674                    t
6675                };
6676                self.spec_qs[j] = q;
6677                all_ids.push(d);
6678                d
6679            } else if sampling {
6680                let mut q = std::mem::take(&mut self.spec_q[j]);
6681                sampler::distribution_into(
6682                    &lg,
6683                    &cfg,
6684                    all_ids,
6685                    &mut self.sampler_scratch,
6686                    self.pool.as_deref(),
6687                    &mut q,
6688                );
6689                let d = sampler::draw(&q, &mut self.rng);
6690                self.spec_q[j] = q;
6691                all_ids.push(d); // the next draft's penalties see this one
6692                d
6693            } else if greedy_pen {
6694                let d = sampler::argmax_penalized(
6695                    &lg,
6696                    &cfg,
6697                    all_ids,
6698                    &mut self.sampler_scratch,
6699                    self.pool.as_deref(),
6700                );
6701                all_ids.push(d);
6702                d
6703            } else {
6704                sampler::argmax(&lg)
6705            };
6706            attention::recycle_buf(&mut lg);
6707            drafts.push(dj);
6708            hx = hj;
6709            spec_stamp("d.pick");
6710        }
6711        all_ids.truncate(base_len);
6712        *drafted += k_spec;
6713        let t_draft = t_round.elapsed();
6714        let sub_draft = subs();
6715        // Verify batch: [t_next, d1 .. d_{k-1}] at next_pos.. — every row's
6716        // logits come back from the graph's own head.
6717        let b = k_spec + 1;
6718        let mut hiddens = vec![0.0f32; b * self.hidden_size];
6719        for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
6720            let e = self.embed_single(t);
6721            hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
6722        }
6723        let positions: Vec<usize> = (next_pos..next_pos + b).collect();
6724        spec_stamp("v.emb");
6725        let (lm_gw, lm_rows) = {
6726            let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6727            (
6728                crate::gpu::GraphW {
6729                    idx: i,
6730                    kind,
6731                    row_scale: rs,
6732                    data: &[],
6733                    prism: crate::gpu::GraphPrismOp::None,
6734                    affine: false,
6735                },
6736                self.weights.lm_head.rows(),
6737            )
6738        };
6739        let mut logits = Vec::new();
6740        let final_norm = self.weights.final_norm.clone();
6741        // Plain greedy on Metal: the b argmaxes come from the device
6742        // (`argmax_rows` after the head) and the 7.9 MB logits plane is
6743        // never read back — the round's decision needs only the ids, and
6744        // the loop top takes the last verified id as `spec_forced`, which
6745        // is exactly what its argmax of the row would give. The full rows
6746        // stay for anything that reads them: sampling, penalties,
6747        // confidence, the verify oracle, the logit dump.
6748        // `CMF_METAL_DEV_ARGMAX=0` keeps the host path.
6749        #[cfg(target_os = "macos")]
6750        let greedy_dev = metal_native
6751            && !sampling
6752            && !greedy_pen
6753            && !self.confidence_on
6754            && self.final_softcap.is_none()
6755            // The host acceptance argmax scans the WHOLE head row
6756            // (`lm_rows`), the sampler's own row only `vocab_size`: they
6757            // coincide exactly when the head has no padding rows, and
6758            // only then is the device argmax (which scores `vocab_size`)
6759            // bit-identical to both.
6760            && self.vocab_size == lm_rows
6761            && std::env::var_os("CMF_METAL_VERIFY_CHECK").is_none()
6762            && std::env::var_os("CMF_LOGIT_DUMP").is_none()
6763            && std::env::var("CMF_METAL_DEV_ARGMAX").as_deref() != Ok("0");
6764        #[cfg(not(target_os = "macos"))]
6765        let greedy_dev = false;
6766        let mut dev_ids: Vec<u32> = Vec::new();
6767        #[cfg(target_os = "macos")]
6768        let verify_outcome = if metal_native {
6769            let lm = self.weights.lm_head.q1_parts()?;
6770            let n_score = self.vocab_size.min(lm_rows);
6771            self.try_batch_graph_metal(
6772                &mut hiddens,
6773                &positions,
6774                b,
6775                Some((lm, &final_norm, &mut logits)),
6776                if greedy_dev {
6777                    Some((n_score, &mut dev_ids))
6778                } else {
6779                    None
6780                },
6781            )
6782        } else {
6783            self.try_batch_graph_wgpu(
6784                &mut hiddens,
6785                &positions,
6786                b,
6787                Some(crate::gpu::SpecTail {
6788                    lm: lm_gw,
6789                    lm_rows,
6790                    final_norm: &final_norm,
6791                    logits_out: &mut logits,
6792                }),
6793            )
6794        };
6795        #[cfg(not(target_os = "macos"))]
6796        let verify_outcome = self.try_batch_graph_wgpu(
6797            &mut hiddens,
6798            &positions,
6799            b,
6800            Some(crate::gpu::SpecTail {
6801                lm: lm_gw,
6802                lm_rows,
6803                final_norm: &final_norm,
6804                logits_out: &mut logits,
6805            }),
6806        );
6807        match verify_outcome {
6808            crate::gpu::BatchGraphOutcome::Completed => {}
6809            crate::gpu::BatchGraphOutcome::Declined => {
6810                // The verifier refused before admission.  Its draft MTP
6811                // rows are still device-resident, so rewind the separate
6812                // mirror before the caller takes the exact one-token path.
6813                m.kv.truncate_last(k_spec);
6814                if !metal_native && !self.rewind_mtp_graph_mirror(next_pos) {
6815                    self.clear_sequence_state();
6816                    self.graph_failed
6817                        .store(true, std::sync::atomic::Ordering::Relaxed);
6818                    self.cancel
6819                        .store(true, std::sync::atomic::Ordering::Relaxed);
6820                    tracing::error!("MTP graph mirror rewind failed after verify decline");
6821                }
6822                return None;
6823            }
6824            crate::gpu::BatchGraphOutcome::Failed => {
6825                // A failed batch may have advanced trunk/GDN state.  Clear
6826                // both mirrors and preserve the terminal outcome rather than
6827                // falling through to stale CPU state.
6828                self.clear_sequence_state();
6829                self.graph_failed
6830                    .store(true, std::sync::atomic::Ordering::Relaxed);
6831                self.cancel
6832                    .store(true, std::sync::atomic::Ordering::Relaxed);
6833                tracing::error!("MTP verify batch graph failed after admission");
6834                return None;
6835            }
6836        }
6837        // `CMF_METAL_VERIFY_CHECK=1`: run the same b tokens through the
6838        // plain per-token path and compare each row's argmax + logits with
6839        // the verify's — the bring-up oracle for the batched graph. The
6840        // plain forwards mutate the CPU state; it is snapshotted and put
6841        // back, and the K/V mirrors re-pointed, before the round goes on.
6842        #[cfg(target_os = "macos")]
6843        if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
6844            let snap: Vec<Vec<f32>> = self
6845                .kv_cache
6846                .layers
6847                .iter()
6848                .map(|l| l.linear_state.clone())
6849                .collect();
6850            let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
6851            let toks: Vec<u32> = std::iter::once(t_next)
6852                .chain(drafts.iter().copied())
6853                .collect();
6854            let want_save = self.graph_want_logits;
6855            self.graph_want_logits = false;
6856            for (i, &t) in toks.iter().enumerate() {
6857                let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
6858                let _ = self.graph_logits.take();
6859                // CMF_SPEC_PLAIN_HIDDEN=1: the next round drafts from the
6860                // plain path's hidden instead of the verify's (an experiment
6861                // on the chain's sensitivity to the half-GEMM noise)
6862                if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
6863                    hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
6864                }
6865                let ref_lg = self.logits_from_hidden(&hi);
6866                let row = &logits[i * lm_rows..(i + 1) * lm_rows];
6867                let ra = sampler::argmax(&ref_lg);
6868                let va = sampler::argmax(row);
6869                let mut md = 0f32;
6870                let mut rms = 0f64;
6871                for j in 0..lm_rows.min(ref_lg.len()) {
6872                    let d = (ref_lg[j] - row[j]).abs();
6873                    md = md.max(d);
6874                    rms += (d as f64) * (d as f64);
6875                }
6876                let mut hd = 0f32;
6877                for j in 0..self.hidden_size {
6878                    hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
6879                }
6880                eprintln!(
6881                    "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
6882                    next_pos + i,
6883                    if ra == va { "OK" } else { "MISMATCH" },
6884                    (rms / lm_rows as f64).sqrt()
6885                );
6886            }
6887            self.graph_want_logits = want_save;
6888            // restore IN PLACE: the pending verify graph wraps these very
6889            // allocations (zero-copy) — replacing the Vec would strand it
6890            for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
6891                if l.linear_state.len() == st.len() {
6892                    l.linear_state.copy_from_slice(&st);
6893                } else {
6894                    l.linear_state = st;
6895                }
6896            }
6897            for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
6898                let extra = l.seq_len.saturating_sub(n0);
6899                if extra > 0 {
6900                    l.truncate_last(extra);
6901                    crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
6902                }
6903            }
6904        }
6905        let t_verify = t_round.elapsed();
6906        let sub_verify = subs();
6907        // Acceptance. Greedy: row i's argmax is the trunk's token after
6908        // input i. Sampling: accept draft i with min(1, p_i/q_i), and on
6909        // the first rejection draw the correction from max(0, p_i − q_i)
6910        // — that token is committed by the loop top as-is (spec_forced).
6911        let mut a = 0usize;
6912        let mut forced: Option<u32> = None;
6913        let ids: Vec<u32> = if sparse {
6914            let mut p = std::mem::take(&mut self.spec_ps);
6915            let mut res = std::mem::take(&mut self.spec_ress);
6916            while a < k_spec {
6917                let ok = sampler::sparse_distribution_into(
6918                    &logits[a * lm_rows..(a + 1) * lm_rows],
6919                    &cfg,
6920                    all_ids,
6921                    &mut self.sampler_scratch,
6922                    self.pool.as_deref(),
6923                    &mut p,
6924                );
6925                if !ok {
6926                    let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
6927                    p.clear();
6928                    p.push((t, 1.0));
6929                }
6930                match sampler::spec_accept_or_correct_sparse(
6931                    &p,
6932                    &self.spec_qs[a],
6933                    drafts[a],
6934                    &mut self.rng,
6935                    &mut res,
6936                ) {
6937                    None => {
6938                        all_ids.push(drafts[a]);
6939                        a += 1;
6940                    }
6941                    Some(c) => {
6942                        forced = Some(c);
6943                        break;
6944                    }
6945                }
6946            }
6947            all_ids.truncate(base_len);
6948            self.spec_ps = p;
6949            self.spec_ress = res;
6950            drafts.clone()
6951        } else if sampling {
6952            let mut p = std::mem::take(&mut self.spec_p);
6953            let mut res = std::mem::take(&mut self.spec_res);
6954            while a < k_spec {
6955                sampler::distribution_into(
6956                    &logits[a * lm_rows..(a + 1) * lm_rows],
6957                    &cfg,
6958                    all_ids,
6959                    &mut self.sampler_scratch,
6960                    self.pool.as_deref(),
6961                    &mut p,
6962                );
6963                match sampler::spec_accept_or_correct(
6964                    &p,
6965                    &self.spec_q[a],
6966                    drafts[a],
6967                    &mut self.rng,
6968                    &mut res,
6969                    self.pool.as_deref(),
6970                ) {
6971                    None => {
6972                        all_ids.push(drafts[a]);
6973                        a += 1;
6974                    }
6975                    Some(c) => {
6976                        forced = Some(c);
6977                        break;
6978                    }
6979                }
6980            }
6981            all_ids.truncate(base_len);
6982            self.spec_p = p;
6983            self.spec_res = res;
6984            // the accepted drafts ARE the verified tokens after inputs 0..a
6985            drafts.clone()
6986        } else if greedy_pen {
6987            // Row i's penalized argmax, penalties over the stream that
6988            // includes the accepted drafts before it — the plain loop's
6989            // exact arithmetic, one pass per row, no working copy.
6990            let mut ids: Vec<u32> = Vec::with_capacity(b);
6991            for i in 0..b {
6992                let t = sampler::argmax_penalized(
6993                    &logits[i * lm_rows..(i + 1) * lm_rows],
6994                    &cfg,
6995                    all_ids,
6996                    &mut self.sampler_scratch,
6997                    self.pool.as_deref(),
6998                );
6999                ids.push(t);
7000                if i < k_spec && t == drafts[i] {
7001                    all_ids.push(t);
7002                } else {
7003                    break;
7004                }
7005            }
7006            all_ids.truncate(base_len);
7007            while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
7008                a += 1;
7009            }
7010            // rows past the first mismatch were never scored; the loop
7011            // top re-samples the last verified row itself.
7012            ids
7013        } else if greedy_dev && dev_ids.len() == b {
7014            let ids = std::mem::take(&mut dev_ids);
7015            while a < k_spec && ids[a] == drafts[a] {
7016                a += 1;
7017            }
7018            ids
7019        } else {
7020            if logits.len() < b * lm_rows {
7021                // the device argmax was asked for and came back short:
7022                // no rows to fall back on — terminal like a failed batch
7023                self.clear_sequence_state();
7024                self.graph_failed
7025                    .store(true, std::sync::atomic::Ordering::Relaxed);
7026                self.cancel
7027                    .store(true, std::sync::atomic::Ordering::Relaxed);
7028                tracing::error!("Metal verify returned neither logits nor argmax ids");
7029                return None;
7030            }
7031            let ids: Vec<u32> = (0..b)
7032                .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
7033                .collect();
7034            while a < k_spec && ids[a] == drafts[a] {
7035                a += 1;
7036            }
7037            ids
7038        };
7039        spec_stamp("acc");
7040        if spec_dbg {
7041            eprintln!(
7042                "spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
7043                drafts, ids
7044            );
7045        }
7046        // CMF_METAL_VERIFY_CHECK=2: the commit oracle — plain-forward the
7047        // a+1 accepted tokens from a snapshot, then diff the replayed GDN
7048        // states and the appended K/V rows against that.
7049        #[cfg(target_os = "macos")]
7050        let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
7051            && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
7052        {
7053            let snap: Vec<Vec<f32>> = self
7054                .kv_cache
7055                .layers
7056                .iter()
7057                .map(|l| l.linear_state.clone())
7058                .collect();
7059            let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
7060            let toks: Vec<u32> = std::iter::once(t_next)
7061                .chain(drafts.iter().copied())
7062                .collect();
7063            let want_save = self.graph_want_logits;
7064            self.graph_want_logits = false;
7065            for (i, &t) in toks.iter().take(a + 1).enumerate() {
7066                let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
7067                let _ = self.graph_logits.take();
7068            }
7069            self.graph_want_logits = want_save;
7070            let plain_states: Vec<Vec<f32>> = self
7071                .kv_cache
7072                .layers
7073                .iter()
7074                .map(|l| l.linear_state.clone())
7075                .collect();
7076            let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7077            let mut rows = Vec::new();
7078            for (li, (l, n0)) in self
7079                .kv_cache
7080                .layers
7081                .iter_mut()
7082                .zip(attn_lens.iter())
7083                .enumerate()
7084            {
7085                let extra = l.seq_len.saturating_sub(*n0);
7086                if extra > 0 {
7087                    let mut kk = Vec::new();
7088                    let mut vv = Vec::new();
7089                    for g in 0..nkv {
7090                        kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7091                        vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7092                    }
7093                    rows.push((li, kk, vv));
7094                    l.truncate_last(extra);
7095                    crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
7096                }
7097            }
7098            for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7099                if l.linear_state.len() == st.len() {
7100                    l.linear_state.copy_from_slice(&st);
7101                } else {
7102                    l.linear_state = st;
7103                }
7104            }
7105            Some((plain_states, rows))
7106        } else {
7107            None
7108        };
7109        let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
7110        // Metal: the MTP cache cut and the round's warm-up SUBMIT come
7111        // BEFORE the trunk commit, so the warm-up's command buffer is
7112        // queued ahead of the GDN replay (second queue) and its wait
7113        // below no longer sits behind the replay — measured: the warm-up's
7114        // wait grew with the accepted count exactly like the replay does
7115        // (8 ms at a=1, 17 ms at a=3, 25 ms at a=5 for ~2 ms of its own
7116        // work). The replay now overlaps the warm-up's readback, the
7117        // round's return and the next draft chain.
7118        #[cfg(target_os = "macos")]
7119        let mut warm_pending: Option<MetalWarmPending> = None;
7120        #[cfg(target_os = "macos")]
7121        if metal_native {
7122            m.kv.truncate_last(k_spec.saturating_sub(1));
7123            if self.mtp_graph_mode == Some(true) {
7124                // the mirror rows below the cut are the CPU rows: re-point,
7125                // no re-upload
7126                crate::gpu_metal::kv_mirror_set_stored(
7127                    self.mtp_kv_id(),
7128                    Self::MTP_LAYER_BASE,
7129                    m.kv.seq_len,
7130                );
7131                if !warm_off && a > 0 {
7132                    let pairs: Vec<(&[f32], u32)> = (0..a)
7133                        .map(|j| {
7134                            (
7135                                &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
7136                                ids[j],
7137                            )
7138                        })
7139                        .collect();
7140                    warm_pending = self.mtp_warm_batch_submit(m, &pairs, next_pos);
7141                }
7142            }
7143            spec_stamp("c.wsub");
7144        }
7145        // a fully-accepted round needs no restore: every input was real.
7146        #[cfg(target_os = "macos")]
7147        if metal_native {
7148            // the Metal verify never wrote its states: the commit replays the
7149            // accepted prefix into the CPU owners and appends the K/V rows
7150            if !self.metal_verify_commit(a) {
7151                self.clear_sequence_state();
7152                self.graph_failed
7153                    .store(true, std::sync::atomic::Ordering::Relaxed);
7154                self.cancel
7155                    .store(true, std::sync::atomic::Ordering::Relaxed);
7156                tracing::error!("Metal verify state/KV handoff failed after admission");
7157                return None;
7158            }
7159            if let Some((plain_states, rows)) = commit_ref {
7160                crate::gpu_metal::queue_fence();
7161                // the commit's replay runs on the second queue: collect it
7162                // before the oracle reads the CPU owners it writes into
7163                let _ = crate::gpu_metal::wait_replay();
7164                let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7165                let mut worst_s = 0f32;
7166                let mut worst_li = 0usize;
7167                for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
7168                    if l.linear_state.len() != ps.len() || ps.is_empty() {
7169                        continue;
7170                    }
7171                    let d = l
7172                        .linear_state
7173                        .iter()
7174                        .zip(ps)
7175                        .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7176                    let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
7177                    let rel = d / n.max(1e-6);
7178                    if rel > worst_s {
7179                        worst_s = rel;
7180                        worst_li = li;
7181                    }
7182                }
7183                let mut worst_k = 0f32;
7184                for (li, kk, vv) in &rows {
7185                    let l = &self.kv_cache.layers[*li];
7186                    let n0 = l.seq_len - (kk.len() / (nkv * hd));
7187                    let mut ck = Vec::new();
7188                    let mut cv = Vec::new();
7189                    for g in 0..nkv {
7190                        ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7191                        cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7192                    }
7193                    if ck.len() == kk.len() {
7194                        let dk = ck
7195                            .iter()
7196                            .zip(kk)
7197                            .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7198                        let dv = cv
7199                            .iter()
7200                            .zip(vv)
7201                            .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7202                        worst_k = worst_k.max(dk).max(dv);
7203                    } else {
7204                        eprintln!(
7205                            "commit-check L{li}: kv row count mismatch {} vs {}",
7206                            ck.len(),
7207                            kk.len()
7208                        );
7209                    }
7210                }
7211                eprintln!(
7212                    "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}"
7213                );
7214            }
7215        }
7216        if !metal_native && a + 1 < b {
7217            let expected_gdn_layers = self.graph_gdn_layer_count();
7218            if expected_gdn_layers > 0
7219                && !crate::gpu::gdn_spec_restore(self.graph_kv_id, a, next_pos, expected_gdn_layers)
7220            {
7221                self.clear_sequence_state();
7222                self.graph_failed
7223                    .store(true, std::sync::atomic::Ordering::Relaxed);
7224                self.cancel
7225                    .store(true, std::sync::atomic::Ordering::Relaxed);
7226                tracing::error!("GDN speculative restore failed after verify");
7227                return None;
7228            }
7229        }
7230        if !metal_native && !self.rewind_trunk_graph_mirrors(next_pos + a + 1) {
7231            // The verify graph committed the full batch, but one of its
7232            // persistent Full-attention mirrors could not be re-pointed to
7233            // the accepted prefix.  Treat that as terminal state failure;
7234            // an exact CPU fallback would otherwise consume stale GDN/KV.
7235            self.clear_sequence_state();
7236            self.graph_failed
7237                .store(true, std::sync::atomic::Ordering::Relaxed);
7238            self.cancel
7239                .store(true, std::sync::atomic::Ordering::Relaxed);
7240            tracing::error!("trunk graph KV rewind failed after speculative verify");
7241            return None;
7242        }
7243        *accepted += a;
7244        // MTP cache: keep the first draft row (its inputs were real), drop
7245        // the chain's, then append the verified pairs the round produced.
7246        // Each of those is a whole MTP block on the per-op path and they
7247        // cost 5.8 ms of a 69 ms round at k=3 — a third of what the
7248        // round's own draft costs. PRICED, and they earn it: skipping
7249        // them (`CMF_SPEC_WARM=0`) drops acceptance from 89% to 81% at
7250        // k=3 and 85% to 74% at k=4, and the tok/s goes nowhere at k=3
7251        // (50.3 against 50.5) and backwards at k=4 (48.1 against 50.1).
7252        // The knob stays so the next person can re-price it after the
7253        // warms are batched instead of assuming either way.
7254        if !metal_native {
7255            // (Metal cut its MTP cache before the trunk commit, above)
7256            m.kv.truncate_last(k_spec.saturating_sub(1));
7257        }
7258        spec_stamp("c.trunc");
7259        if !metal_native
7260            && self.mtp_graph_mode == Some(true)
7261            && !self.rewind_mtp_graph_mirror(next_pos)
7262        {
7263            // The graph draft was admitted, so inability to move its cursor
7264            // back to the real anchor is a state failure, not a capability
7265            // refusal.  Do not warm or continue with a stale mirror.
7266            self.clear_sequence_state();
7267            self.graph_failed
7268                .store(true, std::sync::atomic::Ordering::Relaxed);
7269            self.cancel
7270                .store(true, std::sync::atomic::Ordering::Relaxed);
7271            tracing::error!("MTP graph mirror rewind failed after verify commit");
7272            return None;
7273        }
7274        if !warm_off && a > 0 {
7275            // Graph arm: all accepted pairs in ONE batched run over the
7276            // MTP block; the token graph one by one if the batch declines.
7277            let mut warmed = false;
7278            #[cfg(target_os = "macos")]
7279            if metal_native && self.mtp_graph_mode == Some(true) {
7280                // the batched warm-up was submitted before the trunk
7281                // commit: collect it here; one by one on the token graph
7282                // if it declined (or failed)
7283                warmed = match warm_pending.take() {
7284                    Some(p) => self.mtp_warm_batch_finish(m, p),
7285                    None => false,
7286                };
7287                if !warmed {
7288                    warmed = true;
7289                    for j in 0..a {
7290                        let row =
7291                            hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
7292                        if self
7293                            .mtp_step_metal(m, &row, ids[j], next_pos + j, false)
7294                            .is_none()
7295                        {
7296                            warmed = false;
7297                            break;
7298                        }
7299                    }
7300                }
7301            }
7302            if !warmed && self.mtp_graph_mode != Some(false) && !metal_native {
7303                let rows: Vec<Vec<f32>> = (0..a)
7304                    .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
7305                    .collect();
7306                let pairs: Vec<(&[f32], u32)> = rows
7307                    .iter()
7308                    .zip(ids.iter())
7309                    .map(|(r, &t)| (r.as_slice(), t))
7310                    .collect();
7311                match self.mtp_warm_prefill_pairs(m, &pairs, next_pos) {
7312                    Ok(()) => warmed = true,
7313                    Err(err) => {
7314                        // A warm-up failure after graph admission cannot
7315                        // fall back to `mtp_warm`: the detached CPU cache is
7316                        // not authoritative for the device mirror.  Mark it
7317                        // terminal so the generation caller clears state and
7318                        // returns instead of drafting from stale attention.
7319                        tracing::error!("{err}");
7320                        self.clear_sequence_state();
7321                        self.graph_failed
7322                            .store(true, std::sync::atomic::Ordering::Relaxed);
7323                        self.cancel
7324                            .store(true, std::sync::atomic::Ordering::Relaxed);
7325                        return None;
7326                    }
7327                }
7328            }
7329            if !warmed {
7330                for j in 0..a {
7331                    let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
7332                    let row = row.to_vec();
7333                    self.mtp_warm(m, &row, ids[j], next_pos + j);
7334                }
7335            }
7336        }
7337        // The sampler's contract: logits of the LAST verified position —
7338        // unless a rejected draft already drew the correction, in which
7339        // case the loop top commits that token and samples nothing.
7340        spec_stamp("c.warm");
7341        if let Some(c) = forced {
7342            self.spec_forced = Some(c);
7343            self.graph_logits = None;
7344        } else if greedy_dev && logits.is_empty() {
7345            // the row's argmax IS the token the loop top would pick from
7346            // it (plain greedy, no penalties): commit it as forced
7347            self.spec_forced = Some(ids[a]);
7348            self.graph_logits = None;
7349        } else {
7350            let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
7351            row.resize(self.vocab_size, 0.0);
7352            if let Some(c) = self.final_softcap {
7353                for l in row.iter_mut() {
7354                    *l = c * (*l / c).tanh();
7355                }
7356            }
7357            self.graph_logits = Some(row);
7358        }
7359        let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
7360        spec_stamp("c.row");
7361        // Three phases, not two. The round's wall clock was 4 ms longer
7362        // than draft+verify and the difference had nowhere to be seen:
7363        // the accepted prefix re-runs the MTP block once per token to
7364        // keep the draft head's attention cache warm, and the GDN state
7365        // rolls back on any rejection. Both live here, after the verify.
7366        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7367            let end = subs();
7368            eprintln!(
7369                "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
7370                 commit {:.1} ms/{} sub (accepted {a} of {k_spec}, full-head streak {})",
7371                t_draft.as_secs_f64() * 1e3,
7372                sub_draft - sub0,
7373                (t_verify - t_draft).as_secs_f64() * 1e3,
7374                sub_verify - sub_draft,
7375                (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
7376                end - sub_verify,
7377                self.draft_full_streak,
7378            );
7379        }
7380        // Native Metal's verify tile is flat in b (eight rows for the price
7381        // of one), so a shorter round only forfeits tokens — measured on
7382        // the M4: an essay round at k=2 still verified in 260 ms. The
7383        // adaptation is for cards whose verify grows with the rows.
7384        if k_env.is_none() && !metal_native && !k_capped {
7385            // Slow average and a wide band: a fast one oscillated 2↔3 on
7386            // an essay every other round (measured), which forfeits the
7387            // draft it just paid for.
7388            let f = a as f32 / k_spec.max(1) as f32;
7389            self.spec_acc_ewma += 0.2 * (f - self.spec_acc_ewma);
7390            let mut k_next = k_spec;
7391            if self.spec_acc_ewma >= 0.75 && k_spec < k_max {
7392                k_next = k_spec + 1;
7393            } else if self.spec_acc_ewma < 0.4 && k_spec > 2 {
7394                k_next = k_spec - 1;
7395            }
7396            if k_next != k_spec {
7397                self.spec_acc_ewma = 0.6;
7398                if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7399                    eprintln!("spec-k: {k_spec} → {k_next}");
7400                }
7401            }
7402            self.spec_k_adapt = Some(k_next);
7403        }
7404        spec_stamp("end");
7405        Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
7406    }
7407
7408    /// Micro-benchmark: two single-position forwards vs one fused pair
7409    /// from the current cache state (KV rewound after each probe).
7410    /// Returns (two_singles_ms, fused_pair_ms) per probe, or the (0, 0)
7411    /// sentinel when this model has no pair path to measure — the same
7412    /// answer the o1 arm gives, and the bench prints it the same way.
7413    /// (An architecture that loads its own layers leaves `weights.layers`
7414    /// empty; walking it here was an index panic, found by `bench` on
7415    /// deepseek_v4.)
7416    pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
7417        if !self.pair_supported() {
7418            return (0.0, 0.0);
7419        }
7420        // This is a host-side pair micro-benchmark. It truncates the host KV
7421        // after every probe, so letting the whole-token graph participate
7422        // would leave its device GDN/KV mirror ahead of the next probe and
7423        // poison the process-wide graph verdict before the real generation
7424        // benchmark starts. Keep the existing per-op/GPU arithmetic while
7425        // suppressing only the stateful token graph for this measurement.
7426        let graph_env = std::env::var_os("CMF_GPU_WGPU_GRAPH");
7427        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
7428        let emb1 = self.embed_single(1);
7429        let emb2 = self.embed_single(2);
7430        let pos = self.kv_cache.seq_len();
7431
7432        let t0 = std::time::Instant::now();
7433        for _ in 0..iters {
7434            let _ = self.forward_layers(&emb1, pos, None);
7435            let _ = self.forward_layers(&emb2, pos + 1, None);
7436            for l in &mut self.kv_cache.layers {
7437                l.truncate_last(2);
7438            }
7439        }
7440        let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7441
7442        let t1 = std::time::Instant::now();
7443        for _ in 0..iters {
7444            let _ = self.forward_pair(&emb1, &emb2, pos);
7445            for l in &mut self.kv_cache.layers {
7446                l.truncate_last(2);
7447            }
7448        }
7449        let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7450        match graph_env {
7451            Some(value) => unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", value) },
7452            None => unsafe { std::env::remove_var("CMF_GPU_WGPU_GRAPH") },
7453        }
7454        (singles_ms, pair_ms)
7455    }
7456
7457    /// Fused two-position forward: weight rows are streamed from memory
7458    /// once per layer for both positions. Full layers → fused GQA pair;
7459    /// linear layers → vmf_phase pair (lane 2 state is tentative in the
7460    /// per-layer scratch until the draft is accepted).
7461    /// Whether the fused two-position path covers every layer kind in
7462    /// this model. MLA and KDA run per position (their pair arms are
7463    /// unreachable); the seq prefill falls back to singles for them.
7464    fn pair_supported(&self) -> bool {
7465        // An EMPTY layer stack means the architecture loaded its own and
7466        // this path has nothing to walk. Checking that directly, rather
7467        // than naming each such architecture, is what makes the guard hold
7468        // for the next one: `any()` over no layers is false, so a
7469        // feature-by-feature test says "supported" for a model that has no
7470        // layers here at all.
7471        !self.weights.layers.is_empty()
7472            && self.g3n.is_none()
7473            && !self
7474                .weights
7475                .layers
7476                .iter()
7477                .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
7478    }
7479
7480    fn forward_pair(
7481        &mut self,
7482        emb1: &[f32],
7483        emb2: &[f32],
7484        position: usize,
7485    ) -> (Vec<f32>, Vec<f32>) {
7486        // A two-token prompt starts here, not in the layer walk: decide the
7487        // MiMo placement before the pair's per-op MoE uploads any expert.
7488        self.mimo_moe_prepare();
7489        let mut h1 = emb1.to_vec();
7490        let mut h2 = emb2.to_vec();
7491        let (_nkv, _hd, hs, _rd, eps) = (
7492            self.num_kv_heads,
7493            self.head_dim,
7494            self.hidden_size,
7495            self.rotary_dim,
7496            self.rms_eps,
7497        );
7498        let pool = self.pool.clone();
7499
7500        for li in 0..self.num_layers {
7501            let lw = &self.weights.layers[self.phys_layer(li)];
7502            // Norms into pipeline scratch (4 allocs/layer on the MTP
7503            // decode hot path before this).
7504            inference::rms_norm_into(
7505                &h1,
7506                &lw.input_norm,
7507                self.rms_eps,
7508                self.norm_style,
7509                &mut self.ws.n1,
7510            );
7511            inference::rms_norm_into(
7512                &h2,
7513                &lw.input_norm,
7514                self.rms_eps,
7515                self.norm_style,
7516                &mut self.ws.n2,
7517            );
7518
7519            let (a1, a2) = match &lw.attn {
7520                AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
7521                AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
7522                AttnKind::Bounded(w) => {
7523                    // Two sequential positions of the bounded operator
7524                    // (the ring's causal order is the pair's order).
7525                    let rope = self
7526                        .bounded_rope
7527                        .clone()
7528                        .expect("bounded layer without an installed rotation table");
7529                    let cfg = crate::bounded::BoundedAttnCfg {
7530                        num_heads: self.num_heads,
7531                        num_kv_heads: self.num_kv_heads,
7532                        head_dim: self.head_dim,
7533                        hidden_size: hs,
7534                        scale: self.attn_scale,
7535                        rope: &rope,
7536                        pool: pool.as_deref(),
7537                    };
7538                    let a1 = crate::bounded::bounded_attention(
7539                        &self.ws.n1,
7540                        w,
7541                        &mut self.kv_cache.layers[li],
7542                        &cfg,
7543                    );
7544                    let a2 = crate::bounded::bounded_attention(
7545                        &self.ws.n2,
7546                        w,
7547                        &mut self.kv_cache.layers[li],
7548                        &cfg,
7549                    );
7550                    (a1, a2)
7551                }
7552                AttnKind::Linear(w) => {
7553                    let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
7554                    let layer = &mut self.kv_cache.layers[li];
7555                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7556                    vmf_phase_pair(
7557                        &self.ws.n1,
7558                        &self.ws.n2,
7559                        w,
7560                        &cfg,
7561                        state,
7562                        scratch,
7563                        self.pool.as_deref(),
7564                    )
7565                }
7566                AttnKind::LinearGdn(w) => {
7567                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
7568                    let layer = &mut self.kv_cache.layers[li];
7569                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7570                    gdn_pair(
7571                        &self.ws.n1,
7572                        &self.ws.n2,
7573                        w,
7574                        &cfg,
7575                        state,
7576                        scratch,
7577                        self.pool.as_deref(),
7578                    )
7579                }
7580                AttnKind::ShortConv(w) => {
7581                    let cfg = self
7582                        .short_conv_cfg
7583                        .expect("short-conv layer without short_conv_cfg");
7584                    let layer = &mut self.kv_cache.layers[li];
7585                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7586                    short_conv_pair(
7587                        &self.ws.n1,
7588                        &self.ws.n2,
7589                        w,
7590                        &cfg,
7591                        state,
7592                        scratch,
7593                        self.pool.as_deref(),
7594                    )
7595                }
7596                AttnKind::Full {
7597                    wq,
7598                    wk,
7599                    wv,
7600                    wo,
7601                    q_norm,
7602                    k_norm,
7603                    output_gate,
7604                    softplus_gate,
7605                    bias,
7606                } => {
7607                    let inv_freq_l = self.layer_inv_freq(li);
7608                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
7609                    let cfg = QwenAttnCfg {
7610                        num_heads: self.layer_num_heads(li),
7611                        num_kv_heads: nkv_l,
7612                        head_dim: hd_l,
7613                        hidden_size: hs,
7614                        position,
7615                        inv_freq: &inv_freq_l,
7616                        rotary_dim: rd_l,
7617                        scale: self.attn_scale,
7618                        softcap: self.attn_softcap,
7619                        window: self.layer_window(li),
7620                        v_norm: self.attn_v_norm,
7621                        qk_norm_after_rope: self.qk_norm_after_rope,
7622                        q_norm: q_norm.as_deref(),
7623                        k_norm: k_norm.as_deref(),
7624                        output_gate: *output_gate,
7625                        softplus_gate: softplus_gate
7626                            .as_ref()
7627                            .map(|(gate, per_head)| (gate, *per_head)),
7628                        rope_scale: self.layer_rope_scale(li),
7629                        bias: bias
7630                            .as_ref()
7631                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7632                        rms_eps: eps,
7633                        norm_style: self.norm_style,
7634                        pool: pool.as_deref(),
7635                        v_head_dim: self.layer_v_dim(li),
7636                    };
7637                    attention::qwen_attention_pair(
7638                        &self.ws.n1,
7639                        &self.ws.n2,
7640                        wq,
7641                        wk,
7642                        wv,
7643                        wo,
7644                        &mut self.kv_cache.layers[li],
7645                        &cfg,
7646                    )
7647                }
7648            };
7649            let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
7650                Some(w) => (
7651                    inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
7652                    inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
7653                ),
7654                None => (a1, a2),
7655            };
7656            for i in 0..self.hidden_size {
7657                h1[i] += a1[i];
7658                h2[i] += a2[i];
7659            }
7660            let (mut a1, mut a2) = (a1, a2);
7661            attention::recycle_buf(&mut a1);
7662            attention::recycle_buf(&mut a2);
7663
7664            let lw = &self.weights.layers[self.phys_layer(li)];
7665            inference::rms_norm_into(
7666                &h1,
7667                &lw.post_norm,
7668                self.rms_eps,
7669                self.norm_style,
7670                &mut self.ws.p1,
7671            );
7672            inference::rms_norm_into(
7673                &h2,
7674                &lw.post_norm,
7675                self.rms_eps,
7676                self.norm_style,
7677                &mut self.ws.p2,
7678            );
7679            let (f1, f2) = match &lw.ffn {
7680                // Dual-branch layers need the raw residuals — run the
7681                // two positions through the same fn decode uses.
7682                FfnKind::DenseMoe(dm) => (
7683                    dense_moe_ffn(
7684                        dm,
7685                        &self.ws.p1,
7686                        &h1,
7687                        self.rms_eps,
7688                        self.norm_style,
7689                        self.pool.as_deref(),
7690                    ),
7691                    dense_moe_ffn(
7692                        dm,
7693                        &self.ws.p2,
7694                        &h2,
7695                        self.rms_eps,
7696                        self.norm_style,
7697                        self.pool.as_deref(),
7698                    ),
7699                ),
7700                FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => (
7701                    moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p1, self.pool.as_deref()),
7702                    moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p2, self.pool.as_deref()),
7703                ),
7704                _ => ffn_forward_pair(
7705                    &lw.ffn,
7706                    &self.ws.p1,
7707                    &self.ws.p2,
7708                    self.pool.as_deref(),
7709                    None,
7710                ),
7711            };
7712            let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
7713                Some(w) => (
7714                    inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
7715                    inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
7716                ),
7717                None => (f1, f2),
7718            };
7719            for i in 0..self.hidden_size {
7720                h1[i] += f1[i];
7721                h2[i] += f2[i];
7722            }
7723            let (mut f1, mut f2) = (f1, f2);
7724            attention::recycle_buf(&mut f1);
7725            attention::recycle_buf(&mut f2);
7726            if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
7727                for i in 0..self.hidden_size {
7728                    h1[i] *= sc;
7729                    h2[i] *= sc;
7730                }
7731            }
7732            // Looped Transformer: apply final norm at the end of each loop iteration.
7733            if self.is_loop_end(li) && li + 1 < self.num_layers {
7734                h1 = inference::rms_norm(
7735                    &h1,
7736                    &self.weights.final_norm,
7737                    self.rms_eps,
7738                    self.norm_style,
7739                );
7740                h2 = inference::rms_norm(
7741                    &h2,
7742                    &self.weights.final_norm,
7743                    self.rms_eps,
7744                    self.norm_style,
7745                );
7746            }
7747        }
7748        // Real O(1) prefill pairs may also carry tentative lane-2 recurrent
7749        // state. Commit it before publishing the transition epoch so the
7750        // next serial/device row cannot observe a new attention epoch with an
7751        // old GDN state. Speculative pairs run only when O(1) is inactive and
7752        // retain their existing caller-controlled commit/rollback semantics.
7753        if self.o1_active() {
7754            self.commit_linear_scratch();
7755        }
7756        self.o1_progress();
7757        (h1, h2)
7758    }
7759
7760    /// Commit lane-2 linear states after an accepted draft.
7761    fn commit_linear_scratch(&mut self) {
7762        for layer in &mut self.kv_cache.layers {
7763            if !layer.linear_scratch.is_empty() {
7764                std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
7765                layer.linear_scratch.clear();
7766            }
7767        }
7768    }
7769
7770    /// Forward a full id sequence from a fresh cache and return the
7771    /// logits after the last position (golden-parity harness, bench).
7772    pub fn forward_ids(
7773        &mut self,
7774        ids: &[u32],
7775        task_mask: Option<&TaskMask>,
7776    ) -> Result<Vec<f32>, String> {
7777        #[cfg(target_os = "macos")]
7778        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7779        if ids.is_empty() {
7780            return Err("empty id sequence".to_string());
7781        }
7782        self.clear_sequence_state();
7783        self.check_forward_graph("forward_ids setup", 0)?;
7784        if task_mask.is_none() {
7785            self.o1_begin();
7786        }
7787        let mut hidden = vec![0.0f32; self.hidden_size];
7788        let mut pos = 0usize;
7789        if let Some(b) = &mut self.dsv41 {
7790            let pool = self.pool.clone();
7791            let mut logits = Vec::new();
7792            crate::dsv41::forward_chunk(
7793                &b.0,
7794                &b.1,
7795                &b.2,
7796                &mut b.3,
7797                ids,
7798                0,
7799                pool.as_deref(),
7800                &mut logits,
7801            );
7802            if let Err(err) = self.o1_seal_checked() {
7803                self.clear_sequence_state();
7804                return Err(err);
7805            }
7806            return Ok(logits);
7807        }
7808        // Same routing predicate generation uses. Two reasons it must be
7809        // the same one: (1) a GDN hybrid's recurrent state is GPU-
7810        // resident, and a batched CPU prefill would build it on the host
7811        // only — decode then reads buffers the prefill never wrote;
7812        // (2) bench times THIS function and calls the result "prefill",
7813        // so a different path here reports a number production never
7814        // sees (W2 on 2×5090: 8.7 tok/s reported against 125 real).
7815        if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
7816            // prefill-GEMM in chunks; only the last position's hidden is
7817            // needed. (o1-compatible: the batch path attends per position
7818            // through qwen_attention, which carries the collection hook.)
7819            let chunk = self.prefill_chunk();
7820            let hs = self.hidden_size;
7821            while pos < ids.len() {
7822                let end = (pos + chunk).min(ids.len());
7823                let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
7824                self.check_forward_graph("forward_ids batched prefill", end - 1)?;
7825                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
7826                pos = end;
7827            }
7828        }
7829        // Same guards as generation's prefill — INCLUDING the graph one.
7830        // The CPU pair walk was intercepting positions that the resident
7831        // token graph would have run itself: on a GDN hybrid over wgpu
7832        // that is 89 ms of host forward against 7 ms of device submit,
7833        // and it made prefill look 12× slower than it is (W2 on an RTX
7834        // 5090, ctx 512: 11.2 tok/s with the walk, 136.6 without).
7835        // CMF_PAIR=0 opts out; a model whose layers live outside
7836        // `weights.layers` has no pair walk to take.
7837        if task_mask.is_none()
7838            && !self.graph_prefill_preferred()
7839            && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
7840            && self.pair_supported()
7841        {
7842            while pos + 1 < ids.len() {
7843                let e1 = self.embed_single(ids[pos]);
7844                let e2 = self.embed_single(ids[pos + 1]);
7845                let (_, h2) = self.forward_pair(&e1, &e2, pos);
7846                self.check_forward_graph("forward_ids pair", pos + 1)?;
7847                self.commit_linear_scratch();
7848                hidden = h2;
7849                pos += 2;
7850            }
7851        }
7852        // Resident Embryo graph: the prompt in chunks of one submit each
7853        // (the same device state and logits as the per-position walk).
7854        if task_mask.is_none() && pos == 0 && ids.len() > 1 {
7855            if let Some(lg) = self.embryo_prefill_chunked(ids, 0) {
7856                self.graph_logits = Some(lg);
7857                hidden = vec![0.0; self.hidden_size];
7858                pos = ids.len();
7859            }
7860        }
7861        while pos < ids.len() {
7862            hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
7863            self.check_forward_graph("forward_ids", pos)?;
7864            pos += 1;
7865        }
7866        if let Some(logits) = self.graph_logits.take() {
7867            // Resident stacks already applied final norm and their head in
7868            // the same submit; do not run a second norm/head over the zero
7869            // hidden sentinel returned by forward_layers_span.
7870            if let Err(err) = self.o1_seal_checked() {
7871                self.clear_sequence_state();
7872                return Err(err);
7873            }
7874            return Ok(logits);
7875        }
7876        // Harness contract: after forward_ids the cache is decode-ready —
7877        // under o1 that means sealed (bench measures the seal as part of
7878        // prefill, honestly).
7879        if let Err(err) = self.o1_seal_checked() {
7880            self.clear_sequence_state();
7881            return Err(err);
7882        }
7883        let normed = inference::rms_norm(
7884            &hidden,
7885            &self.weights.final_norm,
7886            self.rms_eps,
7887            self.norm_style,
7888        );
7889        Ok(self.lm_head_forward(&normed))
7890    }
7891
7892    /// Run the V4.1 stack one token at a time and retain logits for every
7893    /// position. This is a diagnostic surface for comparing a converted
7894    /// checkpoint with a tokenwise reference implementation.
7895    #[doc(hidden)]
7896    pub fn dsv41_serial_logits(&mut self, ids: &[u32]) -> Result<Vec<Vec<f32>>, String> {
7897        #[cfg(target_os = "macos")]
7898        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7899        if ids.is_empty() {
7900            return Err("empty id sequence".to_string());
7901        }
7902        self.clear_sequence_state();
7903        self.dsv41
7904            .as_ref()
7905            .ok_or_else(|| "dsv41 serial logits require a DeepSeek-V4.1 model".to_string())?;
7906        self.o1_begin();
7907        let rows = {
7908            let pool = self.pool.clone();
7909            let b = self
7910                .dsv41
7911                .as_mut()
7912                .expect("dsv41 checked above; state cannot change during forward");
7913            let mut rows = Vec::with_capacity(ids.len());
7914            for (position, &id) in ids.iter().enumerate() {
7915                let mut logits = Vec::new();
7916                crate::dsv41::forward_token(
7917                    &b.0,
7918                    &b.1,
7919                    &b.2,
7920                    &mut b.3,
7921                    id,
7922                    position,
7923                    pool.as_deref(),
7924                    &mut logits,
7925                );
7926                rows.push(logits);
7927            }
7928            rows
7929        };
7930        self.o1_seal();
7931        Ok(rows)
7932    }
7933
7934    /// Teacher-forced perplexity over a token sequence (phase-C gate:
7935    /// honest quant comparisons instead of prompt vibes).
7936    ///
7937    /// Attention is EXACT even on a model whose layers are flagged for
7938    /// the O(1) kernel — scoring the backbone is the default on purpose
7939    /// (it is the yardstick). `nll_ids_o1` scores the CONVERTED model.
7940    pub fn ppl_ids(&mut self, ids: &[u32]) -> Result<f64, String> {
7941        let (nll, cnt) = self.nll_ids_from(ids, 0)?;
7942        Ok((nll / cnt.max(1) as f64).exp())
7943    }
7944
7945    /// DTG-MA calibration pass (Patent 2): run `ids` through the model
7946    /// (CPU path, per position) and return each layer's per-neuron
7947    /// activation mass Σ|silu(gate)·up| — the statistic the task-guided
7948    /// FFN mask is derived from.
7949    pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
7950        self.clear_sequence_state();
7951        FFN_PROBE.with(|p| {
7952            *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
7953        });
7954        crate::gpu::cpu_scope(|| {
7955            for (pos, &id) in ids.iter().enumerate() {
7956                let emb = self.embed_single(id);
7957                let _ = self.forward_layers(&emb, pos, None);
7958            }
7959        });
7960        self.clear_sequence_state();
7961        FFN_PROBE
7962            .with(|p| p.borrow_mut().take())
7963            .unwrap_or_default()
7964    }
7965
7966    /// `probe_ffn_mass` over the BATCHED prefill: same accumulator, one
7967    /// sweep instead of one forward per token. What makes the statistic
7968    /// affordable on a 27B.
7969    pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Result<Vec<Vec<f64>>, String> {
7970        if let Err(err) = self.nll_begin() {
7971            // A recorder can be left by a caller that was interrupted before
7972            // this request entered its scoring block.  Consume it even when
7973            // the preflight failure prevents initialization of a new one.
7974            let _ = FFN_PROBE.with(|p| p.borrow_mut().take());
7975            self.nll_end();
7976            return Err(err);
7977        }
7978        FFN_PROBE.with(|p| {
7979            *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
7980        });
7981        let result: Result<(), String> = (|| {
7982            for chunk in ids.chunks(256) {
7983                if chunk.len() < 2 {
7984                    continue;
7985                }
7986                self.nll_ids_masked(chunk, 0, None)?;
7987            }
7988            Ok(())
7989        })();
7990        self.nll_end();
7991        let probe = FFN_PROBE
7992            .with(|p| p.borrow_mut().take())
7993            .unwrap_or_default();
7994        match result {
7995            Ok(()) => Ok(probe),
7996            Err(err) => {
7997                drop(probe);
7998                Err(err)
7999            }
8000        }
8001    }
8002
8003    /// Teacher-forced PPL with a task mask active (sparse execution) —
8004    /// the quality gate for a DTG-MA-masked skill. Sequential per
8005    /// position: the batched prefill path is dense-only.
8006    pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> Result<f64, String> {
8007        self.nll_begin()?;
8008        let result: Result<f64, String> = (|| {
8009            let mut nll = 0f64;
8010            let mut cnt = 0usize;
8011            let mut hidden = vec![0f32; self.hidden_size];
8012            for (pos, &id) in ids.iter().enumerate() {
8013                if pos > 0 {
8014                    inference::rms_norm_into(
8015                        &hidden,
8016                        &self.weights.final_norm,
8017                        self.rms_eps,
8018                        self.norm_style,
8019                        &mut self.ws.n1,
8020                    );
8021                    let mut logits = self.lm_head_forward(&self.ws.n1);
8022                    let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
8023                    let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
8024                    let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
8025                    nll -= p.max(1e-300).ln();
8026                    cnt += 1;
8027                    attention::recycle_buf(&mut logits);
8028                }
8029                let emb = self.embed_single(id);
8030                hidden = self.forward_layers(&emb, pos, Some(mask));
8031                self.nll_check_graph("masked serial forward", pos)?;
8032                // Consume a possible graph logits side channel before the
8033                // next row.  Masked scoring normally disables that route,
8034                // but stale channel state must never survive a request.
8035                let _ = self.graph_logits.take();
8036            }
8037            Ok((nll / cnt.max(1) as f64).exp())
8038        })();
8039        self.nll_end();
8040        result
8041    }
8042
8043    /// Teacher-forced NLL sum + scored-token count over positions
8044    /// `start..len-1`, attention EXACT. Positions below `start` still
8045    /// run — they are the context — they are just not scored, so this
8046    /// pairs with `nll_ids_o1(ids, start)` over the very same tokens.
8047    ///
8048    /// Returning (nll, cnt) rather than a ppl is what lets a windowed
8049    /// caller combine windows before the exp, so every scored token
8050    /// weighs the same regardless of how the windows are cut.
8051    /// `nll_ids_from` with a task mask held active at every position.
8052    ///
8053    /// The batched prefill path does not thread masks, so this walks the
8054    /// per-position forward — slower, but it scores the file exactly the
8055    /// way `run --task` will serve it, which is the point of the gate
8056    /// that calls it. With `None` it defers to the fast path.
8057    /// Masked scoring rides the SAME batched sweep as unmasked scoring —
8058    /// the masked-inference fast path: `prefill_batch_masked` lands the
8059    /// per-visit FFN rows on the activations inside the fused arms. The
8060    /// per-position loop below remains only as the no-batch fallback.
8061    pub fn nll_ids_masked(
8062        &mut self,
8063        ids: &[u32],
8064        start: usize,
8065        task_mask: Option<&TaskMask>,
8066    ) -> Result<(f64, usize), String> {
8067        let task_mask = self.drop_open_mask(task_mask);
8068        self.nll_ids_inner(ids, start, task_mask)
8069    }
8070
8071    pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> Result<(f64, usize), String> {
8072        self.nll_ids_inner(ids, start, None)
8073    }
8074
8075    fn nll_ids_inner(
8076        &mut self,
8077        ids: &[u32],
8078        start: usize,
8079        task_mask: Option<&TaskMask>,
8080    ) -> Result<(f64, usize), String> {
8081        self.nll_begin()?;
8082        let result: Result<(f64, usize), String> = (|| {
8083            let mut nll = 0f64;
8084            let mut cnt = 0usize;
8085            // An unmasked quality run with the resident wgpu graph must score
8086            // the same stateful path used by generation.  The layer-major
8087            // GEMM prefill below is a valid CPU/GEMM oracle, but it seeds
8088            // neither the graph's device GDN state nor its device KV mirrors;
8089            // using it here would silently score a different execution.  Keep
8090            // masked scoring on the exact per-position path as before, and
8091            // let the serial arm below drive the graph-aware scorer.
8092            // Only native Metal has a fused graph lm_head contract.  Vulkan
8093            // and other graph backends may expose hidden state without the
8094            // optional logits side channel; preserve their established CPU
8095            // norm/head fallback instead of turning that valid route into a
8096            // hard missing-logits error.
8097            let (graph_quality, fused_head_quality) = nll_graph_policy(
8098                task_mask.is_none(),
8099                self.graph_prefill_preferred(),
8100                crate::gpu::q1_force(),
8101            );
8102            self.graph_head_required = fused_head_quality;
8103            self.graph_want_logits = fused_head_quality;
8104            #[cfg(target_os = "macos")]
8105            if graph_quality && std::env::var("CMF_METAL_BATCH_NLL").as_deref() != Ok("0") {
8106                match self.nll_batch_metal(ids, start) {
8107                    MetalBatchNllOutcome::Completed(nll, count) => {
8108                        return Ok((nll, count));
8109                    }
8110                    MetalBatchNllOutcome::Declined => {}
8111                    MetalBatchNllOutcome::Failed(err) => return Err(err),
8112                }
8113            }
8114            if self.can_prefill_batched() && !graph_quality {
8115                // prefill-GEMM: layer-major position chunks, lm_head batched
8116                // (254MB lm_head read once per chunk, not per position).
8117                // The layer chunk is large (grouping positions by MoE experts
8118                // wins with size), lm_head in sub-blocks (logit buffer
8119                // 32×vocab ≈ 32MB instead of 128×).
8120                const CHUNK: usize = 128;
8121                const LM_SUB: usize = 32;
8122                let n = ids.len().saturating_sub(1);
8123                let hs = self.hidden_size;
8124                let rows = self.weights.lm_head.rows();
8125                let mut pos = 0usize;
8126                let state_trace = std::env::var("CMF_STATE_TRACE").is_ok();
8127                while pos < n {
8128                    let end = (pos + CHUNK).min(n);
8129                    let bsz = end - pos;
8130                    let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
8131                    self.nll_check_graph("batched prefill", pos)?;
8132                    if state_trace && end % 256 == 0 {
8133                        self.trace_recurrent_state(end);
8134                    }
8135                    let mut k0 = 0usize;
8136                    while k0 < bsz {
8137                        let k1 = (k0 + LM_SUB).min(bsz);
8138                        let sb = k1 - k0;
8139                        // Sub-block entirely below the scored range: the KV
8140                        // it just built is all this pass needed from it.
8141                        if pos + k1 <= start {
8142                            k0 = k1;
8143                            continue;
8144                        }
8145                        let mut normed = vec![0.0f32; sb * hs];
8146                        for k in 0..sb {
8147                            let r = inference::rms_norm(
8148                                &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
8149                                &self.weights.final_norm,
8150                                self.rms_eps,
8151                                self.norm_style,
8152                            );
8153                            normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
8154                        }
8155                        let mut logits = vec![0.0f32; sb * rows];
8156                        self.weights
8157                            .lm_head
8158                            .matmat(&normed, sb, &mut logits, self.pool.as_deref());
8159                        for k in 0..sb {
8160                            if pos + k0 + k < start {
8161                                continue;
8162                            }
8163                            self.nll_check_graph("batched score row", pos + k0 + k)?;
8164                            let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
8165                            if let Some(mu) = self.logit_multiplier {
8166                                for v in lg.iter_mut() {
8167                                    *v *= mu;
8168                                }
8169                            }
8170                            // Gemma-class final-logit soft-capping: the
8171                            // decode paths apply it; scoring must too, or
8172                            // the uncapped softmax misprices every token.
8173                            if let Some(c) = self.final_softcap {
8174                                for v in lg.iter_mut() {
8175                                    *v = c * (*v / c).tanh();
8176                                }
8177                            }
8178                            // Cortiq Embryo hierarchical head: same correction
8179                            // the decode path applies (lm_head_forward).
8180                            if let Some(cm) = self.head_clusters.clone() {
8181                                self.hierarchical_head_logprobs(
8182                                    &normed[k * hs..(k + 1) * hs],
8183                                    &cm,
8184                                    lg,
8185                                );
8186                            }
8187                            let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
8188                            let target = ids[pos + k0 + k + 1] as usize;
8189                            let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8190                            let lse: f64 = lg
8191                                .iter()
8192                                .map(|&v| ((v - max) as f64).exp())
8193                                .sum::<f64>()
8194                                .ln()
8195                                + max as f64;
8196                            nll += lse - lg[target] as f64;
8197                            cnt += 1;
8198                            if std::env::var("CMF_PPL_TRACE").is_ok() {
8199                                let top = lg
8200                                    .iter()
8201                                    .enumerate()
8202                                    .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8203                                    .map(|(i, _)| i)
8204                                    .unwrap_or(0);
8205                                eprintln!(
8206                                    "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
8207                                    pos + k0 + k,
8208                                    target,
8209                                    lse - lg[target] as f64,
8210                                    top,
8211                                    lg[target],
8212                                    lg[top]
8213                                );
8214                            }
8215                        }
8216                        k0 = k1;
8217                    }
8218                    pos = end;
8219                }
8220                return Ok((nll, cnt));
8221            }
8222            for pos in 0..ids.len().saturating_sub(1) {
8223                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
8224                self.nll_check_graph("serial forward", pos)?;
8225                // Architectures whose head lives inside their own stack return
8226                // the logits out of band and a zero hidden — DeepSeek-V4 folds
8227                // its hyper-connection copies between the last layer and the
8228                // norm, so it cannot hand back a vector this loop could use.
8229                // Scoring the zeros gave a perplexity of exactly the vocabulary
8230                // size, which is a uniform distribution reported as a
8231                // measurement. `generate` already reads this channel.
8232                let out_of_band = self.graph_logits.take();
8233                if self.graph_head_required && out_of_band.is_none() {
8234                    METAL_GRAPH_HEAD_MISS.fetch_add(
8235                        1,
8236                        std::sync::atomic::Ordering::Relaxed,
8237                    );
8238                    return Err(format!(
8239                        "fused Metal graph head did not complete at NLL position {pos}"
8240                    ));
8241                }
8242                if pos < start {
8243                    continue;
8244                }
8245                let logits = match out_of_band {
8246                    Some(lg) => lg,
8247                    None => {
8248                        let normed = inference::rms_norm(
8249                            &hidden,
8250                            &self.weights.final_norm,
8251                            self.rms_eps,
8252                            self.norm_style,
8253                        );
8254                        // lm_head_forward applies the final-logit softcap itself
8255                        // — capping again here double-squashed gemma-class
8256                        // logits (tanh∘tanh) and reported a flattered ppl.
8257                        self.lm_head_forward(&normed)
8258                    }
8259                };
8260                let target = ids[pos + 1] as usize;
8261                let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8262                let lse: f64 = logits
8263                    .iter()
8264                    .map(|&v| ((v - max) as f64).exp())
8265                    .sum::<f64>()
8266                    .ln()
8267                    + max as f64;
8268                let tok_nll = lse - logits[target] as f64;
8269                if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8270                    let top = logits
8271                        .iter()
8272                        .enumerate()
8273                        .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8274                        .map(|(i, _)| i)
8275                        .unwrap_or(0);
8276                    eprintln!(
8277                        "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8278                        logits[target], logits[top]
8279                    );
8280                }
8281                nll += tok_nll;
8282                cnt += 1;
8283            }
8284            Ok((nll, cnt))
8285        })();
8286        self.nll_end();
8287        result
8288    }
8289
8290    /// Score one post-layer hidden with the same final norm/head path used by
8291    /// decode. Keeping this in one helper is important for the production
8292    /// batch scorer: its rows stop before the final norm, just like the
8293    /// per-position O(1) path below.
8294    fn nll_from_hidden(&mut self, hidden: &[f32], target: u32, pos: usize) -> f64 {
8295        let normed = inference::rms_norm(
8296            hidden,
8297            &self.weights.final_norm,
8298            self.rms_eps,
8299            self.norm_style,
8300        );
8301        // lm_head_forward applies the final-logit softcap itself — capping
8302        // again here double-squashed gemma-class logits in earlier scorers.
8303        let mut logits = self.lm_head_forward(&normed);
8304        let target = target as usize;
8305        let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8306        let lse: f64 = logits
8307            .iter()
8308            .map(|&v| ((v - max) as f64).exp())
8309            .sum::<f64>()
8310            .ln()
8311            + max as f64;
8312        let tok_nll = lse - logits[target] as f64;
8313        if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8314            let top = logits
8315                .iter()
8316                .enumerate()
8317                .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8318                .map(|(i, _)| i)
8319                .unwrap_or(0);
8320            eprintln!(
8321                "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8322                logits[target], logits[top]
8323            );
8324        }
8325        attention::recycle_buf(&mut logits);
8326        tok_nll
8327    }
8328
8329    /// `CMF_STATE_TRACE`: per-layer magnitude of the recurrent record at
8330    /// position `pos` — the whole `linear_state` (vmf: S then the conv
8331    /// ring; GDN: conv ring then S), its recurrent S part alone, the
8332    /// bounded ring, and the last per-position KV row count.  The tool
8333    /// that separated the ~4k perplexity cliff of the 500-step exports
8334    /// (a state that keeps climbing past the trained window) from a
8335    /// runtime boundary; one line per layer, `STATE pos=… layer=…`.
8336    fn trace_recurrent_state(&self, pos: usize) {
8337        let stats = |v: &[f32]| -> (f64, f64) {
8338            if v.is_empty() {
8339                return (0.0, 0.0);
8340            }
8341            let (mut ss, mut mx) = (0f64, 0f64);
8342            for &x in v {
8343                ss += (x as f64) * (x as f64);
8344                mx = mx.max((x as f64).abs());
8345            }
8346            ((ss / v.len() as f64).sqrt(), mx)
8347        };
8348        for (li, l) in self.kv_cache.layers.iter().enumerate() {
8349            let lw = &self.weights.layers[self.phys_layer(li)];
8350            let (kind, s_len) = match &lw.attn {
8351                AttnKind::Linear(_) => (
8352                    "vmf",
8353                    self.vmf_cfg.map(|c| c.state_len()).unwrap_or(0),
8354                ),
8355                AttnKind::LinearGdn(_) => (
8356                    "gdn",
8357                    self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0),
8358                ),
8359                AttnKind::Bounded(_) => ("bounded", 0),
8360                AttnKind::Full { .. } => ("full", 0),
8361                _ => ("other", 0),
8362            };
8363            let (rms, max) = stats(&l.linear_state);
8364            let s_part = if kind == "vmf" {
8365                &l.linear_state[..s_len.min(l.linear_state.len())]
8366            } else if kind == "gdn" {
8367                let ring = l.linear_state.len().saturating_sub(s_len).min(l.linear_state.len());
8368                &l.linear_state[ring..]
8369            } else {
8370                &l.linear_state[..0]
8371            };
8372            let (s_rms, s_max) = stats(s_part);
8373            let (ring_rms, ring_len) = match &l.bounded {
8374                Some(b) => (stats(&b.ring_k).0, b.ring_k.len()),
8375                None => (0.0, 0),
8376            };
8377            eprintln!(
8378                "STATE pos={pos} layer={li} kind={kind} state_len={} rms={rms:.5} max={max:.4} \
8379                 S_rms={s_rms:.5} S_max={s_max:.4} ring_k_rms={ring_rms:.5} ring_len={ring_len} kv_rows={}",
8380                l.linear_state.len(),
8381                l.seq_len
8382            );
8383        }
8384    }
8385
8386    /// Teacher-forced NLL of the CONVERTED model: the O(1) Nyström path
8387    /// is ACTIVE over the scored positions. Returns `Ok((nll sum, scored
8388    /// count))` over `prefill..len-1` and surfaces a post-mutation batch
8389    /// failure instead of returning a partial score.
8390    ///
8391    /// Runtime discipline, deliberately NOT the matrix probe's: the
8392    /// requested prefix plus any required deferred lead-in run the exact
8393    /// prompt pass — that pass is what freezes the landmarks and M — and
8394    /// every post-seal scored position goes through `NystromState::step()`,
8395    /// the same code decode runs.
8396    /// So the landmarks are PREFILL-frozen (what ships), not
8397    /// full-sequence oracles (what the published probe measured). When the
8398    /// requested prefix is shorter than the bounded transition, rows in the
8399    /// exact lead-in are still scored so the shifted target range is stable.
8400    ///
8401    /// Pair with `nll_ids_from(ids, prefill)` for the exact baseline
8402    /// over the identical token set — that ratio is the honest one.
8403    pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> Result<(f64, usize), String> {
8404        // This scorer consumes host hiddens, so never request the optional
8405        // token-graph lm_head side channel. `nll_begin` also consumes a
8406        // prior graph failure and clears only the cancel bit that failure
8407        // raised, leaving a caller-owned cancellation observable.
8408        self.nll_begin()?;
8409        let requested_prefix = (prefill > 0).then_some(prefill);
8410        self.o1_begin_with_prefix(requested_prefix);
8411        let n = ids.len().saturating_sub(1);
8412        let requested_start = prefill.min(n);
8413        // The exact prefix must reach the deferred boundary before a
8414        // collecting layer can convert. Rows between the requested start and
8415        // that boundary remain part of the public NLL range and are scored
8416        // from the same hidden pass below.
8417        let exact_end = if self.o1_active() {
8418            match requested_prefix {
8419                Some(requested) => self.o1_effective_boundary(requested),
8420                None => self
8421                    .o1_cfg
8422                    .as_ref()
8423                    .and_then(|c| crate::nystrom::o1_deferred_boundary(c.w, c.sink)),
8424            }
8425            .unwrap_or(requested_start)
8426            .min(n)
8427        } else {
8428            requested_start
8429        };
8430        let mut nll = 0f64;
8431        let mut cnt = 0usize;
8432
8433        // Exact prompt pass over ids[..exact_end]: the seal consumes its
8434        // q/k/v. Rows at or after requested_start are scored here when the
8435        // bounded lead-in is longer than the caller's requested prefix.
8436        let mut pos = 0usize;
8437        if self.can_prefill_batched() {
8438            const CHUNK: usize = 128;
8439            while pos < exact_end {
8440                let end = (pos + CHUNK).min(exact_end);
8441                let hiddens = self.prefill_batch(&ids[pos..end], pos);
8442                if self
8443                    .graph_failed
8444                    .swap(false, std::sync::atomic::Ordering::Relaxed)
8445                {
8446                    self.cancel
8447                        .store(false, std::sync::atomic::Ordering::Relaxed);
8448                    self.nll_end();
8449                    return Err("GPU graph failed during O(1) NLL prefix".into());
8450                }
8451                for row in 0..end - pos {
8452                    let score_pos = pos + row;
8453                    if score_pos >= requested_start && score_pos < n {
8454                        nll += self.nll_from_hidden(
8455                            &hiddens[row * self.hidden_size..(row + 1) * self.hidden_size],
8456                            ids[score_pos + 1],
8457                            score_pos,
8458                        );
8459                        cnt += 1;
8460                    }
8461                }
8462                pos = end;
8463            }
8464        } else {
8465            while pos < exact_end {
8466                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8467                if self
8468                    .graph_failed
8469                    .swap(false, std::sync::atomic::Ordering::Relaxed)
8470                {
8471                    self.cancel
8472                        .store(false, std::sync::atomic::Ordering::Relaxed);
8473                    self.nll_end();
8474                    return Err("GPU graph failed during O(1) NLL prefix".into());
8475                }
8476                if pos >= requested_start && pos < n {
8477                    nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8478                    cnt += 1;
8479                }
8480                pos += 1;
8481            }
8482        }
8483        self.o1_seal_checked().map_err(|err| {
8484            self.nll_end();
8485            err
8486        })?;
8487
8488        // Reuse the production whole-token batch graph for the post-seal
8489        // suffix when the caller explicitly enabled both routes. This is a
8490        // teacher-forced scorer, so every row is ids[pos] and its target is
8491        // ids[pos + 1]; no speculative tail or rollback state is involved.
8492        // A first Declined is safe to handle with the established serial O(1)
8493        // path. Once a chunk completes, however, the device recurrent state
8494        // owns the sequence and a later decline must be terminal rather than
8495        // falling back to stale CPU state.
8496        let batch_k = std::env::var("CMF_BATCH_K")
8497            .ok()
8498            .and_then(|v| v.parse::<usize>().ok())
8499            .unwrap_or(0);
8500        let batch_admitted = batch_k > 0
8501            && self.can_prefill_batched()
8502            && self.o1_active()
8503            && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
8504            && (0..self.num_layers).all(|li| {
8505                let cache = &self.kv_cache.layers[self.phys_layer(li)];
8506                cache.o1.is_none() || cache.o1_views().is_some()
8507            });
8508        if std::env::var("CMF_GRAPH_PROF").is_ok() {
8509            eprintln!(
8510                "nll-batch: phase=post-seal admission={} requested_k={} scored_rows={}",
8511                batch_admitted,
8512                batch_k,
8513                n.saturating_sub(exact_end),
8514            );
8515        }
8516        let mut batch_completed = false;
8517        if batch_admitted && exact_end < n {
8518            let hs = self.hidden_size;
8519            let mut batch_pos = exact_end;
8520            while batch_pos < n {
8521                let end = (batch_pos + batch_k).min(n);
8522                let bk = end - batch_pos;
8523                let mut hiddens = vec![0.0f32; bk * hs];
8524                for (row, &id) in ids[batch_pos..end].iter().enumerate() {
8525                    hiddens[row * hs..(row + 1) * hs].copy_from_slice(&self.embed_single(id));
8526                }
8527                let positions: Vec<usize> = (batch_pos..end).collect();
8528                let t_batch = std::time::Instant::now();
8529                let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
8530                if std::env::var("CMF_GRAPH_PROF").is_ok() {
8531                    let ms = t_batch.elapsed().as_secs_f64() * 1000.0;
8532                    eprintln!(
8533                        "nll-batch: phase=post-seal mode=o1 k={bk} pos={}..{} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
8534                        batch_pos,
8535                        end.saturating_sub(1),
8536                        bk as f64 / (ms / 1000.0),
8537                    );
8538                }
8539                if let Err(err) = self.nll_check_graph("batch graph", batch_pos) {
8540                    self.nll_end();
8541                    return Err(err);
8542                }
8543                match outcome {
8544                    crate::gpu::BatchGraphOutcome::Completed => {
8545                        batch_completed = true;
8546                        for row in 0..bk {
8547                            nll += self.nll_from_hidden(
8548                                &hiddens[row * hs..(row + 1) * hs],
8549                                ids[batch_pos + row + 1],
8550                                batch_pos + row,
8551                            );
8552                            cnt += 1;
8553                        }
8554                        batch_pos = end;
8555                    }
8556                    crate::gpu::BatchGraphOutcome::Declined => {
8557                        if batch_completed {
8558                            self.nll_end();
8559                            return Err(format!(
8560                                "O(1) NLL batch declined after completed chunk at position {batch_pos}"
8561                            ));
8562                        }
8563                        break;
8564                    }
8565                    crate::gpu::BatchGraphOutcome::Failed => {
8566                        self.nll_end();
8567                        return Err(format!(
8568                            "O(1) NLL batch graph failed after admission at position {batch_pos}"
8569                        ));
8570                    }
8571                }
8572            }
8573            if batch_completed && cnt == n.saturating_sub(requested_start) {
8574                self.nll_end();
8575                return Ok((nll, cnt));
8576            }
8577        }
8578
8579        // Serial O(1) fallback/reference. It is intentionally retained when
8580        // batch admission declines before mutation; callers must label this
8581        // CMF_BATCH_K=0/per-position path separately from the production
8582        // whole-token batch route.
8583        for pos in exact_end..n {
8584            let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8585            if self
8586                .graph_failed
8587                .swap(false, std::sync::atomic::Ordering::Relaxed)
8588            {
8589                self.cancel
8590                    .store(false, std::sync::atomic::Ordering::Relaxed);
8591                self.nll_end();
8592                return Err(format!(
8593                    "GPU graph failed during O(1) NLL serial scoring at position {pos}"
8594                ));
8595            }
8596            nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8597            cnt += 1;
8598        }
8599        self.nll_end();
8600        Ok((nll, cnt))
8601    }
8602
8603    /// Teacher-forced calibration data (B1): for each position, whether the
8604    /// argmax equals the actual next token, and the top-1 softmax prob
8605    /// (top-1 probability) under EACH temperature in `temps` — all from ONE forward
8606    /// pass (argmax/correctness are temperature-invariant; only p_max
8607    /// reshapes). Feeds `cortiq calibrate` (reliability/ECE + temperature
8608    /// fit): is the model's confidence a true property, or does it need a
8609    /// measured scaling?
8610    pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
8611        self.clear_sequence_state();
8612        let n = ids.len().saturating_sub(1);
8613        let mut correct = Vec::with_capacity(n);
8614        let mut pmax = Vec::with_capacity(n);
8615        for pos in 0..n {
8616            let emb = self.embed_single(ids[pos]);
8617            let hidden = self.forward_layers(&emb, pos, None);
8618            let logits = if let Some(logits) = self.graph_logits.take() {
8619                logits
8620            } else {
8621                let normed = inference::rms_norm(
8622                    &hidden,
8623                    &self.weights.final_norm,
8624                    self.rms_eps,
8625                    self.norm_style,
8626                );
8627                // lm_head_forward applies the final-logit softcap itself —
8628                // capping again here double-squashed gemma-class logits
8629                // (tanh∘tanh) and reported a flattered ppl.
8630                self.lm_head_forward(&normed)
8631            };
8632            let target = ids[pos + 1] as usize;
8633            let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
8634            for (i, &v) in logits.iter().enumerate() {
8635                if v > mval {
8636                    mval = v;
8637                    amax = i;
8638                }
8639            }
8640            correct.push(amax == target);
8641            let row: Vec<f32> = temps
8642                .iter()
8643                .map(|&t| {
8644                    let tt = t.max(1e-3);
8645                    let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
8646                    1.0 / s.max(1e-12) // numerator at the max is exp(0)=1
8647                })
8648                .collect();
8649            pmax.push(row);
8650        }
8651        self.clear_sequence_state();
8652        (correct, pmax)
8653    }
8654
8655    /// Teacher-forced PPL with the dynamic router driving per-window
8656    /// skill switches (VMF experiment №2 measurement). Sequential (φ
8657    /// must update per token), returns (ppl, switch_count). The router
8658    /// must be enabled (`enable_dynamic_routing`); else this equals
8659    /// plain `ppl_ids`. The active skill when scoring token t shapes the
8660    /// logits for t+1 — on-policy over the held-out text itself.
8661    pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> Result<(f64, usize), String> {
8662        if self.dyn_router.is_none() {
8663            return Ok((self.ppl_ids(ids)?, 0));
8664        }
8665        self.nll_begin()?;
8666        let saved_active = self.dyn_active;
8667        let mut router = self
8668            .dyn_router
8669            .take()
8670            .ok_or_else(|| "dynamic router disappeared before PPL scoring".to_string())?;
8671        router.reset();
8672        self.dyn_phi_seen = 0;
8673        let _ = self.set_active_skill(None);
8674
8675        let result: Result<(f64, usize), String> = (|| {
8676            let mut nll = 0f64;
8677            let mut cnt = 0usize;
8678            for pos in 0..ids.len().saturating_sub(1) {
8679                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8680                self.nll_check_graph("dynamic serial forward", pos)?;
8681                let out_of_band = self.graph_logits.take();
8682                let mut logits = match out_of_band {
8683                    Some(lg) => lg,
8684                    None => {
8685                        let normed = inference::rms_norm(
8686                            &hidden,
8687                            &self.weights.final_norm,
8688                            self.rms_eps,
8689                            self.norm_style,
8690                        );
8691                        // lm_head_forward applies the final-logit softcap itself —
8692                        // capping again here double-squashed gemma-class logits
8693                        // and reported a flattered ppl.
8694                        self.lm_head_forward(&normed)
8695                    }
8696                };
8697                let target = ids[pos + 1] as usize;
8698                let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8699                let lse: f64 = logits
8700                    .iter()
8701                    .map(|&v| ((v - max) as f64).exp())
8702                    .sum::<f64>()
8703                    .ln()
8704                    + max as f64;
8705                let tok_nll = lse - logits[target] as f64;
8706                if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8707                    let top = logits
8708                        .iter()
8709                        .enumerate()
8710                        .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8711                        .map(|(i, _)| i)
8712                        .unwrap_or(0);
8713                    eprintln!(
8714                        "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8715                        logits[target], logits[top]
8716                    );
8717                }
8718                nll += tok_nll;
8719                cnt += 1;
8720                attention::recycle_buf(&mut logits);
8721                // Route on the evolving phi (drives the NEXT token's skill).
8722                let phi = self.dyn_phi_ema.clone();
8723                if let Some(new_active) = router.step(&phi, pos) {
8724                    let _ = self.set_active_skill(new_active);
8725                }
8726            }
8727            Ok(((nll / cnt.max(1) as f64).exp(), router.switches.len()))
8728        })();
8729
8730        // Restore the detached router and the active overlay on both success
8731        // and failure. The scoring state is cleared independently below.
8732        let _ = self.set_active_skill(saved_active);
8733        self.dyn_router = Some(router);
8734        self.nll_end();
8735        result
8736    }
8737
8738    /// Routing probe φ (spec §9): mean-pooled hidden after `layer`.
8739    pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
8740        self.clear_sequence_state();
8741        let mut acc = vec![0f32; self.hidden_size];
8742        for (pos, &id) in ids.iter().enumerate() {
8743            let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8744            for (a, v) in acc.iter_mut().zip(&h) {
8745                *a += v;
8746            }
8747        }
8748        let n = ids.len().max(1) as f32;
8749        for a in acc.iter_mut() {
8750            *a /= n;
8751        }
8752        self.clear_sequence_state();
8753        acc
8754    }
8755
8756    /// Router-v2 φ probe (spec §9.4, `phi.pool = "span_mean"`): the hidden
8757    /// AFTER `layer` — the same per-position walk and the same quantity as
8758    /// [`Self::probe_phi`] — averaged over the positions in `span` only
8759    /// (the user text between the template's prefix and suffix ids), NOT
8760    /// unit-normalized (the decision normalizes). The walk stops at
8761    /// `span.end`: causality makes the later positions irrelevant, so the
8762    /// result is bit-identical to probing `ids[..span.end]`. An empty span
8763    /// gives the zero vector, which the decision treats as degenerate.
8764    ///
8765    /// Every sequence state is reset before and after — the host KV/ring/
8766    /// recurrent state, the reuse keys (`kv_history`, `kv_prefix`) and the
8767    /// device graph's sequence — so run it on a pipeline that does not
8768    /// also serve a conversation (its prefix reuse would be lost).
8769    pub fn probe_phi_span(
8770        &mut self,
8771        ids: &[u32],
8772        layer: usize,
8773        span: std::ops::Range<usize>,
8774    ) -> Vec<f32> {
8775        #[cfg(target_os = "macos")]
8776        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8777        let end = span.end.min(ids.len());
8778        let start = span.start.min(end);
8779        let reset = |p: &mut Self| p.clear_sequence_state();
8780        reset(self);
8781        let mut acc = vec![0f32; self.hidden_size];
8782        for (pos, &id) in ids[..end].iter().enumerate() {
8783            let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8784            if pos >= start {
8785                for (a, v) in acc.iter_mut().zip(&h) {
8786                    *a += v;
8787                }
8788            }
8789        }
8790        let n = end - start;
8791        if n > 0 {
8792            let n = n as f32;
8793            for a in acc.iter_mut() {
8794                *a /= n;
8795            }
8796        }
8797        reset(self);
8798        acc
8799    }
8800
8801    /// One decode step of the current sequence: forward `token` at
8802    /// `position` (the cache holds positions `[0, position)`, e.g. after
8803    /// [`Self::forward_ids`]) and return the next-token logits — the same
8804    /// forward and head the generation loop runs (resident-graph logits
8805    /// when the graph ran, final norm + lm_head otherwise). The logit-dump
8806    /// tools drive greedy decoding with it so every position is observable.
8807    pub fn decode_step_logits(&mut self, token: u32, position: usize) -> Vec<f32> {
8808        #[cfg(target_os = "macos")]
8809        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8810        self.graph_logits = None;
8811        let hidden = self.forward_layers(&self.embed_single(token), position, None);
8812        if let Some(logits) = self.graph_logits.take() {
8813            return logits;
8814        }
8815        inference::rms_norm_into(
8816            &hidden,
8817            &self.weights.final_norm,
8818            self.rms_eps,
8819            self.norm_style,
8820            &mut self.ws.n1,
8821        );
8822        self.lm_head_forward(&self.ws.n1)
8823    }
8824
8825    /// Layer-major batched prefill (prefill-GEMM): full-attention —
8826    /// per-position with the existing operators (KV grows naturally,
8827    /// causality preserved), GDN projections / FFN / MoE — batched
8828    /// (a weight row is read from DRAM once per chunk, not per
8829    /// position). Returns the hidden of all positions [b × hidden].
8830    fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
8831        self.prefill_batch_masked(ids, start_pos, None)
8832    }
8833
8834    /// `prefill_batch` with a task mask honored on the dense-FFN panels
8835    /// (the masked-inference fast path: full fused compute, mask lands on
8836    /// the activations). The whole-chunk GPU graph is skipped for masked
8837    /// layers by the callers' arms; the per-GEMM device paths stay in
8838    /// play because the zeroing happens on the host between them.
8839    fn prefill_batch_masked(
8840        &mut self,
8841        ids: &[u32],
8842        start_pos: usize,
8843        task_mask: Option<&TaskMask>,
8844    ) -> Vec<f32> {
8845        self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
8846    }
8847
8848    /// One prompt chunk through the whole stack, post-stack rows out (no
8849    /// final norm) — the ingest generation uses, shared by scoring and
8850    /// `forward_ids` so they measure the same execution: the batched wgpu
8851    /// graph's device prefix plus the host's batched walk for the rest when
8852    /// `batch_prefix_prefill` holds and the graph admits the chunk, else
8853    /// the host's chunked prefill. Err only when a graph that had mutated
8854    /// device state failed.
8855    fn prefill_rows(
8856        &mut self,
8857        ids: &[u32],
8858        pos: usize,
8859        task_mask: Option<&TaskMask>,
8860    ) -> Result<Vec<f32>, String> {
8861        self.prefill_input_rows(PrefillIn::Ids(ids), pos, task_mask)
8862    }
8863
8864    fn prefill_input_rows(
8865        &mut self,
8866        input: PrefillIn<'_>,
8867        pos: usize,
8868        task_mask: Option<&TaskMask>,
8869    ) -> Result<Vec<f32>, String> {
8870        self.mimo_moe_prepare();
8871        let hs = self.hidden_size;
8872        let bk = match input {
8873            PrefillIn::Ids(ids) => ids.len(),
8874            PrefillIn::Hidden(rows) => rows.len() / hs,
8875        };
8876        #[cfg(not(target_os = "macos"))]
8877        if task_mask.is_none()
8878            && !self.o1_active()
8879            && bk > 1
8880            && (self.batch_prefix_prefill()
8881                || (self.verify_exact_moe
8882                    && crate::gpu::enabled_here()
8883                    && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)))
8884        {
8885            let mut hiddens = match input {
8886                PrefillIn::Hidden(rows) => rows.to_vec(),
8887                PrefillIn::Ids(ids) => ids.iter().flat_map(|&id| self.embed_single(id)).collect(),
8888            };
8889            let positions: Vec<usize> = (pos..pos + bk).collect();
8890            let mut run = 0usize;
8891            match self.try_batch_graph_wgpu_prefix(
8892                &mut hiddens,
8893                &positions,
8894                bk,
8895                None,
8896                Some(&mut run),
8897            ) {
8898                crate::gpu::BatchGraphOutcome::Completed => {
8899                    let out = if run < self.num_layers {
8900                        self.prefill_batch_span(
8901                            PrefillIn::Hidden(&hiddens),
8902                            pos,
8903                            None,
8904                            run,
8905                            self.num_layers,
8906                        )
8907                    } else {
8908                        hiddens
8909                    };
8910                    return if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
8911                        Err("MiMo attention graph failed after admission".into())
8912                    } else { Ok(out) };
8913                }
8914                crate::gpu::BatchGraphOutcome::Failed => {
8915                    return Err("batched prefix prefill failed after admission".into());
8916                }
8917                crate::gpu::BatchGraphOutcome::Declined => {
8918                    // Rows an earlier chunk left on the device only.
8919                    #[cfg(feature = "gpu")]
8920                    self.pull_lagging_host_kv(0, self.num_layers, pos);
8921                }
8922            }
8923        }
8924        let out = self.prefill_batch_span(input, pos, task_mask, 0, usize::MAX);
8925        if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
8926            Err("batch tail graph failed after admission".into())
8927        } else { Ok(out) }
8928    }
8929
8930    /// The layer-major batched walk over a layer span [from..upto_excl):
8931    /// the whole prefill machinery (chunk graph, batched attends, GEMM
8932    /// panels) for a PARTIAL stack — the network split's prefill rides
8933    /// the same canon as the local one. Input is token ids (embeds
8934    /// itself, coordinator side) or ready boundary hiddens (worker side).
8935    fn prefill_batch_span(
8936        &mut self,
8937        input: PrefillIn<'_>,
8938        start_pos: usize,
8939        task_mask: Option<&TaskMask>,
8940        from: usize,
8941        upto_excl: usize,
8942    ) -> Vec<f32> {
8943        let hs = self.hidden_size;
8944        let b = match input {
8945            PrefillIn::Ids(ids) => ids.len(),
8946            PrefillIn::Hidden(hb) => hb.len() / hs,
8947        };
8948        let upto_excl = upto_excl.min(self.num_layers);
8949        // The CPU embed is deferred: when the chunk graph takes the run
8950        // from layer 0 it gathers the embeddings on the device instead.
8951        // A hidden input is ready by definition.
8952        let mut h: Vec<f32>;
8953        let mut h_ready;
8954        match input {
8955            PrefillIn::Ids(_) => {
8956                h = vec![0.0; b * hs];
8957                h_ready = false;
8958            }
8959            PrefillIn::Hidden(hb) => {
8960                h = hb.to_vec();
8961                h_ready = true;
8962            }
8963        }
8964        let fill_h = |h: &mut Vec<f32>, me: &Self| {
8965            if let PrefillIn::Ids(ids) = input {
8966                for (bi, &id) in ids.iter().enumerate() {
8967                    let e = me.embed_single(id);
8968                    h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
8969                }
8970                if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
8971                    if let Ok(t) = tp.parse::<usize>() {
8972                        if t >= start_pos && t < start_pos + ids.len() {
8973                            let bi = t - start_pos;
8974                            let row = &h[bi * hs..(bi + 1) * hs];
8975                            let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
8976                            eprintln!(
8977                                "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
8978                                ids[bi],
8979                                row[0],
8980                                row[1],
8981                                ids.len(),
8982                                &ids[..ids.len().min(8)]
8983                            );
8984                        }
8985                    }
8986                }
8987            }
8988        };
8989        let (_nkv, _hd, _rd, eps) = (
8990            self.num_kv_heads,
8991            self.head_dim,
8992            self.rotary_dim,
8993            self.rms_eps,
8994        );
8995        let pool = self.pool.clone();
8996        let norm_style = self.norm_style;
8997        self.mimo_moe_prepare();
8998        let automatic_gpu_prefix = self.automatic_gpu_prefix();
8999
9000        #[cfg(target_os = "macos")]
9001        let mut chunk_skip_until = 0usize;
9002        for li in from..upto_excl {
9003            let _capacity_tail = automatic_gpu_prefix
9004                .filter(|&prefix| {
9005                    li >= prefix && !(self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false))
9006                })
9007                .map(|_| crate::gpu::enter_cpu_scope());
9008            crate::gpu::set_layer(li as i64); // layer-split GPU/CPU
9009            // GPU chunk graph (default-on under CMF_GPU=1): a run of
9010            // consecutive eligible layers for the whole chunk in ONE
9011            // Metal submission — norm, QKV, RoPE with fused mirror
9012            // append, causal attend, O, FFN, hidden device-resident
9013            // across the run. Any refusal falls through to the CPU path.
9014            #[cfg(target_os = "macos")]
9015            if task_mask.is_none() {
9016                if li < chunk_skip_until {
9017                    continue;
9018                }
9019                // Device-side embedding needs a q8_row embedding matrix;
9020                // with any other layout the CPU fills `h` first and the
9021                // graph starts from a ready hidden (refusing the whole
9022                // run over the embedding alone kept q4t models — the
9023                // whole Nanbeige/Bonsai class — on the CPU prefill).
9024                if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
9025                    fill_h(&mut h, self);
9026                    h_ready = true;
9027                }
9028                let ids_for_embed = match input {
9029                    PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
9030                    PrefillIn::Hidden(_) => None,
9031                };
9032                let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
9033                if end > li {
9034                    h_ready = true;
9035                    chunk_skip_until = end;
9036                    // Looped Transformer: the graph stopped at a loop
9037                    // boundary — apply final norm before the next iteration.
9038                    if self.is_loop_end(end - 1) && end < self.num_layers {
9039                        for bi in 0..b {
9040                            let normed = inference::rms_norm(
9041                                &h[bi * hs..(bi + 1) * hs],
9042                                &self.weights.final_norm,
9043                                eps,
9044                                norm_style,
9045                            );
9046                            h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9047                        }
9048                    }
9049                    continue;
9050                }
9051            }
9052            if !h_ready {
9053                fill_h(&mut h, self);
9054                h_ready = true;
9055            }
9056            if task_mask.is_none() && self.verify_exact_moe {
9057                let positions: Vec<_> = (start_pos..start_pos + b).collect();
9058                match self.mimo_graph_layer_rows(li, &mut h, &positions) {
9059                    crate::gpu::BatchGraphOutcome::Completed => continue,
9060                    crate::gpu::BatchGraphOutcome::Failed => return h,
9061                    crate::gpu::BatchGraphOutcome::Declined => {},
9062                }
9063            }
9064            #[cfg(feature = "gpu")]
9065            self.pull_lagging_host_kv(li, li + 1, start_pos);
9066            let lw = &self.weights.layers[self.phys_layer(li)];
9067            // ── attention ──
9068            match &lw.attn {
9069                AttnKind::Kda(w) => {
9070                    // Projections batched, recurrence sequential.
9071                    let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
9072                    let mut normed = vec![0.0f32; b * hs];
9073                    for bi in 0..b {
9074                        inference::rms_norm_into(
9075                            &h[bi * hs..(bi + 1) * hs],
9076                            &lw.input_norm,
9077                            eps,
9078                            norm_style,
9079                            &mut normed[bi * hs..(bi + 1) * hs],
9080                        );
9081                    }
9082                    let attn = crate::linear_core::kda_forward_batch(
9083                        &normed,
9084                        b,
9085                        w,
9086                        &cfg,
9087                        &mut self.kv_cache.layers[li].linear_state,
9088                        pool.as_deref(),
9089                    );
9090                    for (dst, &a) in h.iter_mut().zip(&attn) {
9091                        *dst += a;
9092                    }
9093                }
9094                AttnKind::LinearGdn(w) => {
9095                    // Projections batched, recurrence sequential.
9096                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
9097                    let mut normed = vec![0.0f32; b * hs];
9098                    for bi in 0..b {
9099                        let r = inference::rms_norm(
9100                            &h[bi * hs..(bi + 1) * hs],
9101                            &lw.input_norm,
9102                            eps,
9103                            norm_style,
9104                        );
9105                        normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9106                    }
9107                    let attn = crate::linear_core::gdn_forward_batch(
9108                        &normed,
9109                        b,
9110                        w,
9111                        &cfg,
9112                        &mut self.kv_cache.layers[li].linear_state,
9113                        pool.as_deref(),
9114                    );
9115                    for (dst, &a) in h.iter_mut().zip(&attn) {
9116                        *dst += a;
9117                    }
9118                }
9119                AttnKind::ShortConv(w) => {
9120                    // Projections batched over the chunk; the conv walks the
9121                    // contiguous positions in order (same ring as decode).
9122                    let cfg = self
9123                        .short_conv_cfg
9124                        .expect("short-conv layer without short_conv_cfg");
9125                    let mut normed = vec![0.0f32; b * hs];
9126                    for bi in 0..b {
9127                        inference::rms_norm_into(
9128                            &h[bi * hs..(bi + 1) * hs],
9129                            &lw.input_norm,
9130                            eps,
9131                            norm_style,
9132                            &mut normed[bi * hs..(bi + 1) * hs],
9133                        );
9134                    }
9135                    let attn = short_conv_forward_batch(
9136                        &normed,
9137                        b,
9138                        w,
9139                        &cfg,
9140                        &mut self.kv_cache.layers[li].linear_state,
9141                        pool.as_deref(),
9142                    );
9143                    for (dst, &a) in h.iter_mut().zip(&attn) {
9144                        *dst += a;
9145                    }
9146                }
9147                AttnKind::Mla(w) => {
9148                    // Per-position prefill (correctness first; latent
9149                    // batching is a later optimization).
9150                    let inv_freq_l = self.layer_inv_freq(li);
9151                    let rs = self.layer_rope_scale(li);
9152                    let mut normed = vec![0.0f32; hs];
9153                    for bi in 0..b {
9154                        inference::rms_norm_into(
9155                            &h[bi * hs..(bi + 1) * hs],
9156                            &lw.input_norm,
9157                            eps,
9158                            norm_style,
9159                            &mut normed,
9160                        );
9161                        let ao = mla_attention(
9162                            w,
9163                            &normed,
9164                            &mut self.kv_cache.layers[li],
9165                            start_pos + bi,
9166                            &inv_freq_l,
9167                            rs,
9168                            eps,
9169                            pool.as_deref(),
9170                        );
9171                        for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
9172                            *dst += a;
9173                        }
9174                    }
9175                }
9176                AttnKind::Full {
9177                    wq,
9178                    wk,
9179                    wv,
9180                    wo,
9181                    q_norm,
9182                    k_norm,
9183                    output_gate,
9184                    softplus_gate,
9185                    bias,
9186                } => {
9187                    // Chunk-GEMM QKV/O; per-position causal attention
9188                    // inside (roadmap §3 P0 — full-attention prefill no
9189                    // longer re-reads the projection weights b times).
9190                    let mut normed = vec![0.0f32; b * hs];
9191                    for bi in 0..b {
9192                        inference::rms_norm_into(
9193                            &h[bi * hs..(bi + 1) * hs],
9194                            &lw.input_norm,
9195                            eps,
9196                            norm_style,
9197                            &mut normed[bi * hs..(bi + 1) * hs],
9198                        );
9199                    }
9200                    let inv_freq_l = self.layer_inv_freq(li);
9201                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
9202                    let cfg = QwenAttnCfg {
9203                        num_heads: self.layer_num_heads(li),
9204                        num_kv_heads: nkv_l,
9205                        head_dim: hd_l,
9206                        hidden_size: hs,
9207                        position: start_pos,
9208                        inv_freq: &inv_freq_l,
9209                        rotary_dim: rd_l,
9210                        scale: self.attn_scale,
9211                        softcap: self.attn_softcap,
9212                        window: self.layer_window(li),
9213                        v_norm: self.attn_v_norm,
9214                        qk_norm_after_rope: self.qk_norm_after_rope,
9215                        q_norm: q_norm.as_deref(),
9216                        k_norm: k_norm.as_deref(),
9217                        output_gate: *output_gate,
9218                        softplus_gate: softplus_gate
9219                            .as_ref()
9220                            .map(|(gate, per_head)| (gate, *per_head)),
9221                        rope_scale: self.layer_rope_scale(li),
9222                        bias: bias
9223                            .as_ref()
9224                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
9225                        rms_eps: eps,
9226                        norm_style,
9227                        pool: pool.as_deref(),
9228                        v_head_dim: self.layer_v_dim(li),
9229                    };
9230                    let mut attn = attention::qwen_attention_batch(
9231                        &normed,
9232                        b,
9233                        wq,
9234                        wk,
9235                        wv,
9236                        wo,
9237                        &mut self.kv_cache.layers[li],
9238                        &cfg,
9239                    );
9240                    if let Some(w) = &lw.attn_out_norm {
9241                        for bi in 0..b {
9242                            inference::rms_norm_into(
9243                                &attn[bi * hs..(bi + 1) * hs],
9244                                w,
9245                                eps,
9246                                norm_style,
9247                                &mut normed[bi * hs..(bi + 1) * hs],
9248                            );
9249                        }
9250                        attn.copy_from_slice(&normed);
9251                    }
9252                    for (dst, &a) in h.iter_mut().zip(&attn) {
9253                        *dst += a;
9254                    }
9255                }
9256                AttnKind::Bounded(w) => {
9257                    // Chunk-GEMM projections, the bounded operator per
9258                    // position over ring + chunk — never a growing KV.
9259                    let mut normed = vec![0.0f32; b * hs];
9260                    for bi in 0..b {
9261                        inference::rms_norm_into(
9262                            &h[bi * hs..(bi + 1) * hs],
9263                            &lw.input_norm,
9264                            eps,
9265                            norm_style,
9266                            &mut normed[bi * hs..(bi + 1) * hs],
9267                        );
9268                    }
9269                    let rope = self
9270                        .bounded_rope
9271                        .clone()
9272                        .expect("bounded layer without an installed rotation table");
9273                    let cfg = crate::bounded::BoundedAttnCfg {
9274                        num_heads: self.num_heads,
9275                        num_kv_heads: self.num_kv_heads,
9276                        head_dim: self.head_dim,
9277                        hidden_size: hs,
9278                        scale: self.attn_scale,
9279                        rope: &rope,
9280                        pool: pool.as_deref(),
9281                    };
9282                    let mut attn = crate::bounded::bounded_attention_batch(
9283                        &normed,
9284                        b,
9285                        w,
9286                        &mut self.kv_cache.layers[li],
9287                        &cfg,
9288                    );
9289                    if let Some(wn) = &lw.attn_out_norm {
9290                        for bi in 0..b {
9291                            inference::rms_norm_into(
9292                                &attn[bi * hs..(bi + 1) * hs],
9293                                wn,
9294                                eps,
9295                                norm_style,
9296                                &mut normed[bi * hs..(bi + 1) * hs],
9297                            );
9298                        }
9299                        attn.copy_from_slice(&normed);
9300                    }
9301                    for (dst, &a) in h.iter_mut().zip(&attn) {
9302                        *dst += a;
9303                    }
9304                    attention::recycle_buf(&mut attn);
9305                }
9306                AttnKind::Linear(w) => {
9307                    for bi in 0..b {
9308                        let normed = inference::rms_norm(
9309                            &h[bi * hs..(bi + 1) * hs],
9310                            &lw.input_norm,
9311                            eps,
9312                            norm_style,
9313                        );
9314                        vmf_phase_forward(
9315                            &normed,
9316                            w,
9317                            &self.vmf_cfg.expect("linear layer without vmf_cfg"),
9318                            &mut self.kv_cache.layers[li].linear_state,
9319                            pool.as_deref(),
9320                        )
9321                        .iter()
9322                        .enumerate()
9323                        .for_each(|(i, &a)| h[bi * hs + i] += a);
9324                    }
9325                }
9326            }
9327
9328            // ── FFN batched ──
9329            let lw = &self.weights.layers[self.phys_layer(li)];
9330            let mut post = vec![0.0f32; b * hs];
9331            for bi in 0..b {
9332                let r =
9333                    inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
9334                post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9335            }
9336            // A restrictive per-visit FFN row lands on the activations
9337            // inside the dense arm; an all-open row costs nothing.
9338            let mask_row = task_mask
9339                .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
9340                .and_then(|m| m.ffn_masks.get(li))
9341                .map(|v| v.as_slice());
9342            let mut ffn = match &lw.ffn {
9343                FfnKind::Dense(d) if !d.segs.is_empty() => {
9344                    tube_ffn(d, &post, b, pool.as_deref(), mask_row)
9345                }
9346                FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
9347                FfnKind::Moe(m) if self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false) => {
9348                    moe_ffn_banked_rows(&mut self.mimo_moe, li, m, &post, b, hs, pool.as_deref())
9349                }
9350                FfnKind::Moe(m) if self.verify_exact_moe => {
9351                    moe_ffn_rows_exact(m, &post, b, hs, pool.as_deref())
9352                }
9353                // Keep prompt expert panels off the projection arena and
9354                // use their routes to prime the model-wide bank.
9355                FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => {
9356                    let before = m.stats.borrow().clone();
9357                    let out = crate::gpu::cpu_scope(|| {
9358                        moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None)
9359                    });
9360                    self.mimo_moe.prime(li, m, &before);
9361                    out
9362                }
9363                FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
9364                // Dual-branch layers run per position (the expert branch
9365                // reads the raw residual — nothing to batch yet).
9366                FfnKind::DenseMoe(dm) => {
9367                    let mut out = vec![0.0f32; b * hs];
9368                    for bi in 0..b {
9369                        let r = dense_moe_ffn(
9370                            dm,
9371                            &post[bi * hs..(bi + 1) * hs],
9372                            &h[bi * hs..(bi + 1) * hs],
9373                            eps,
9374                            norm_style,
9375                            pool.as_deref(),
9376                        );
9377                        out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9378                    }
9379                    out
9380                }
9381            };
9382            if let Some(w) = &lw.ffn_out_norm {
9383                for bi in 0..b {
9384                    inference::rms_norm_into(
9385                        &ffn[bi * hs..(bi + 1) * hs],
9386                        w,
9387                        eps,
9388                        norm_style,
9389                        &mut post[bi * hs..(bi + 1) * hs],
9390                    );
9391                }
9392                ffn.copy_from_slice(&post);
9393            }
9394            for (dst, &f) in h.iter_mut().zip(&ffn) {
9395                *dst += f;
9396            }
9397            if let Some(sc) = lw.layer_scale {
9398                for v in h.iter_mut() {
9399                    *v *= sc;
9400                }
9401            }
9402            // CMF_LAYER_DUMP: every position's hidden after layer li.
9403            if self.layer_dump.is_some() {
9404                for bi in 0..b {
9405                    self.dump_layer_row(start_pos + bi, li, &h[bi * hs..(bi + 1) * hs]);
9406                }
9407            }
9408            if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9409                if let Ok(t) = tp.parse::<usize>() {
9410                    if t >= start_pos && t < start_pos + b {
9411                        let bi = t - start_pos;
9412                        let row = &h[bi * hs..(bi + 1) * hs];
9413                        let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9414                        eprintln!(
9415                            "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
9416                            row[0], row[1]
9417                        );
9418                    }
9419                }
9420            }
9421            // CMF_DEBUG_LAYERS=1: per-layer hidden-state health of the
9422            // LAST prompt position — the knife for "which layer type
9423            // breaks first" on a new architecture.
9424            if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
9425                let row = &h[(b - 1) * hs..b * hs];
9426                let rms =
9427                    (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
9428                let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
9429                eprintln!(
9430                    "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
9431                    match &self.weights.layers[self.phys_layer(li)].attn {
9432                        AttnKind::LinearGdn(_) => "gdn",
9433                        AttnKind::Linear(_) => "vmf",
9434                        AttnKind::ShortConv(_) => "conv",
9435                        _ => "attn",
9436                    },
9437                    match &lw.ffn {
9438                        FfnKind::Moe(_) => "moe",
9439                        FfnKind::Dense(_) => "dense",
9440                        FfnKind::DenseMoe(_) => "dense+moe",
9441                    },
9442                );
9443            }
9444            // Looped Transformer: apply final norm at the end of each loop iteration.
9445            if self.is_loop_end(li) && li + 1 < self.num_layers {
9446                for bi in 0..b {
9447                    let normed = inference::rms_norm(
9448                        &h[bi * hs..(bi + 1) * hs],
9449                        &self.weights.final_norm,
9450                        eps,
9451                        norm_style,
9452                    );
9453                    h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9454                }
9455            }
9456            if std::env::var("CMF_TRACE_H").is_ok() {
9457                let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
9458                let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
9459                eprintln!(
9460                    "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
9461                    lw.layer_scale
9462                );
9463            }
9464        }
9465        crate::gpu::set_layer(-1); // lm_head/final ops outside layer-split
9466        // A batched span owns a complete set of positions. Publish any
9467        // collecting→sealed transition only after every layer has finished;
9468        // callers that cross into serial/device work must see the new epoch
9469        // before this function returns.
9470        self.o1_progress();
9471        h
9472    }
9473
9474    /// Embed a single token.
9475    fn embed_single(&self, id: u32) -> Vec<f32> {
9476        let mut out = vec![0.0f32; self.hidden_size];
9477        if (id as usize) < self.weights.embed_tokens.rows() {
9478            self.weights.embed_tokens.row_f32(id as usize, &mut out);
9479        }
9480        if self.embed_multiplier != 1.0 {
9481            for v in out.iter_mut() {
9482                *v *= self.embed_multiplier;
9483            }
9484        }
9485        // DeepSeek-V4's hash layers route by TOKEN ID, so the id has to
9486        // reach the forward. It rides in slot 0 (the forward re-reads the
9487        // real embedding itself from the table).
9488        if self.dsv4.is_some()
9489            || self.dsv41.is_some()
9490            || self.qwen4_exp.is_some()
9491        {
9492            let mut v = vec![0.0f32; self.hidden_size.max(1)];
9493            v[0] = id as f32;
9494            return v;
9495        }
9496        // Gemma-3n: the per-layer-embedding half needs the token ID, so
9497        // it rides appended to the embedding; the g3n forward splits it.
9498        if let Some(b) = &self.g3n {
9499            return b.0.extend_embedding(id, &out, self.pool.as_deref());
9500        }
9501        out
9502    }
9503
9504    /// A run of consecutive prefill layers on the GPU for the whole
9505    /// chunk (default-on under CMF_GPU=1; CMF_GPU_CHUNK=0 disables).
9506    /// Eligibility per layer: q8_row weights, plain full attention
9507    /// (no output gate), F32 KV, no o1/masks/gemma extras. Returns the
9508    /// first layer index NOT processed (== `li0` when the run is empty).
9509    #[cfg(target_os = "macos")]
9510    fn chunk_run_gpu(
9511        &mut self,
9512        li0: usize,
9513        h: &mut [f32],
9514        b: usize,
9515        pos0: usize,
9516        embed_ids: Option<&[u32]>,
9517        cap: usize,
9518    ) -> usize {
9519        // (The old streaming attend needed a depth bound at ~1k; the
9520        // GEMM attention scales like the CPU path and lifted it.)
9521        // CMF_GPU_CHUNK=0 disables the graph.
9522        if !crate::gpu::enabled_here()
9523            || std::env::var("CMF_GPU_CHUNK")
9524                .map(|v| v == "0")
9525                .unwrap_or(false)
9526            || b < 32
9527            || self.swa.is_some()
9528            || self.global_attn.is_some()
9529            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
9530            || self.graph_attn_decline_reason().is_some()
9531            // Collection owns the exact Q trace and boundary conversion;
9532            // this chunk graph appends dense KV without feeding that trace.
9533            || self.o1_active()
9534            || self.attn_v_norm
9535            || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
9536        {
9537            return li0;
9538        }
9539        let Some(model) = self.model.clone() else {
9540            return li0;
9541        };
9542        let inv_freq = self.inv_freq.clone();
9543        let (nh, nkv, hd, hs) = (
9544            self.num_heads,
9545            self.num_kv_heads,
9546            self.head_dim,
9547            self.hidden_size,
9548        );
9549        // Collect the longest run of consecutive eligible layers.
9550        // Looped Transformer: stop at the loop boundary so the CPU can
9551        // apply loop_final_norm between iterations.
9552        let loop_end = if self.loop_final_norm {
9553            ((li0 / self.physical_layers) + 1) * self.physical_layers
9554        } else {
9555            self.num_layers
9556        };
9557        let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
9558        let mut stored_at: Vec<usize> = Vec::new();
9559        for li in li0..self.num_layers.min(loop_end).min(cap) {
9560            let lw = &self.weights.layers[self.phys_layer(li)];
9561            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
9562                break;
9563            }
9564            let AttnKind::Full {
9565                wq,
9566                wk,
9567                wv,
9568                wo,
9569                q_norm,
9570                k_norm,
9571                output_gate: false,
9572                softplus_gate: None,
9573                bias,
9574            } = &lw.attn
9575            else {
9576                break;
9577            };
9578            let FfnKind::Dense(d) = &lw.ffn else { break };
9579            if d.act != Act::Silu || !d.segs.is_empty() {
9580                break;
9581            }
9582            // q8_row (row_scale populated), or q4_tiled / q4tp (row_scale
9583            // empty — their scales are in the payload). Mixing across the
9584            // seven projections of one layer is fine; the encoder branches
9585            // per weight on the tensor's dtype. Anything else refuses.
9586            fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
9587                t.q8_row_parts()
9588                    .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9589                    .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9590            }
9591            let parts = (
9592                cw(wq),
9593                cw(wk),
9594                cw(wv),
9595                cw(wo),
9596                cw(&d.gate_proj),
9597                cw(&d.up_proj),
9598                cw(&d.down_proj),
9599            );
9600            let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
9601            else {
9602                break;
9603            };
9604            let layer = &self.kv_cache.layers[li];
9605            if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
9606                break;
9607            }
9608            stored_at.push(layer.head_len(0));
9609            layers.push(crate::gpu_metal::ChunkLayer {
9610                model: &model,
9611                kv_id: self.graph_kv_id,
9612                layer: li,
9613                wq: pq,
9614                wk: pk,
9615                wv: pv,
9616                wo: po,
9617                gate: pg,
9618                up: pu,
9619                down: pd,
9620                input_norm: &lw.input_norm,
9621                post_norm: &lw.post_norm,
9622                bias: bias
9623                    .as_ref()
9624                    .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
9625                q_norm: q_norm.as_deref(),
9626                k_norm: k_norm.as_deref(),
9627                inv_freq: &inv_freq,
9628                rd: self.rotary_dim,
9629                nh,
9630                nkv,
9631                hd,
9632                hs,
9633                inter: d.gate_proj.rows(),
9634                gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
9635                late_qk_norm: self.qk_norm_after_rope,
9636                eps: self.rms_eps as f32,
9637            });
9638        }
9639        if layers.is_empty() {
9640            return li0;
9641        }
9642        let row = nkv * hd;
9643        let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
9644            .iter()
9645            .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
9646            .collect();
9647        let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
9648        for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
9649            let li = layers[i].layer;
9650            let layer = &self.kv_cache.layers[li];
9651            io.push(crate::gpu_metal::ChunkIo {
9652                cpu_stored: stored_at[i],
9653                cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
9654                cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
9655                out_k: ok,
9656                out_v: ov,
9657                imp: oi,
9658            });
9659        }
9660        let n_run = layers.len();
9661        let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
9662        // Device-side embedding when the run starts the model and the
9663        // embedding matrix is q8_row-mapped.
9664        let ep = embed_ids.and_then(|ids| {
9665            self.weights
9666                .embed_tokens
9667                .q8_row_parts()
9668                .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
9669                    idx,
9670                    rows,
9671                    row_scale: rs,
9672                    ids,
9673                    mult: self.embed_multiplier,
9674                })
9675        });
9676        if embed_ids.is_some() && ep.is_none() {
9677            return li0;
9678        }
9679        if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
9680            return li0;
9681        }
9682        drop(io);
9683        drop(layers);
9684        // CPU caches stay the owners of record: append the chunk rows
9685        // and bank the importance masses per layer.
9686        for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
9687            let li = li0 + i;
9688            let layer = &mut self.kv_cache.layers[li];
9689            for bi in 0..b {
9690                layer.append(
9691                    &ok[bi * row..(bi + 1) * row],
9692                    &ov[bi * row..(bi + 1) * row],
9693                    &[],
9694                );
9695            }
9696            layer.accumulate_imp(oi);
9697        }
9698        last
9699    }
9700
9701    /// Is layer `li` a sliding-window (local-RoPE) layer? Gemma-3:
9702    /// every `pattern`-th layer is global, the rest are local.
9703    fn layer_is_local(&self, li: usize) -> bool {
9704        if let Some(layers) = &self.sliding_layers {
9705            return layers.get(li).copied().unwrap_or(false);
9706        }
9707        match self.swa {
9708            Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
9709            None => false,
9710        }
9711    }
9712
9713    /// The RoPE table for layer `li` (local layers may have their own;
9714    /// Gemma-4 global layers use the proportional padded table).
9715    fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
9716        if self.layer_is_local(li) {
9717            if let Some(f) = &self.inv_freq_local {
9718                return f.clone();
9719            }
9720        } else if let Some(f) = &self.inv_freq_global {
9721            return f.clone();
9722        }
9723        self.inv_freq.clone()
9724    }
9725
9726    /// The attend window for layer `li` (None = full context).
9727    fn layer_window(&self, li: usize) -> Option<usize> {
9728        self.swa
9729            .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
9730    }
9731
9732    fn layer_num_heads(&self, li: usize) -> usize {
9733        self.attention_heads_per_layer
9734            .as_ref()
9735            .and_then(|v| v.get(li).copied())
9736            .unwrap_or(self.num_heads)
9737    }
9738
9739    fn layer_rope_scale(&self, li: usize) -> f32 {
9740        if self.layer_is_local(li) {
9741            self.rope_scale_local
9742        } else {
9743            self.rope_scale
9744        }
9745    }
9746
9747    /// Attention geometry of layer `li`: (num_kv_heads, head_dim,
9748    /// rotary_dim). Gemma-4 global layers override all three.
9749    fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
9750        if !self.layer_is_local(li) {
9751            if let Some((ghd, gkv)) = self.global_attn {
9752                return (gkv, ghd, ghd);
9753            }
9754        }
9755        (
9756            self.layer_num_kv_heads(li),
9757            self.head_dim,
9758            if self.layer_is_local(li) {
9759                self.rotary_dim_local.unwrap_or(self.rotary_dim)
9760            } else {
9761                self.rotary_dim
9762            },
9763        )
9764    }
9765
9766    /// KV heads of layer `li` (virtual index): the per-layer count when the
9767    /// model has one (MiMo-V2), else the uniform `num_kv_heads`.
9768    fn layer_num_kv_heads(&self, li: usize) -> usize {
9769        self.kv_heads_per_layer
9770            .as_ref()
9771            .and_then(|v| v.get(self.phys_layer(li)).copied())
9772            .unwrap_or(self.num_kv_heads)
9773    }
9774
9775    /// V head width of layer `li` (≤ its head_dim).
9776    fn layer_v_dim(&self, li: usize) -> usize {
9777        let (_, hd, _) = self.layer_geom(li);
9778        self.v_head_dim.unwrap_or(hd).min(hd)
9779    }
9780
9781    /// Install a per-layer KV geometry: KV heads per PHYSICAL layer and/or
9782    /// a V head width narrower than `head_dim` (MiMo-V2). Validates it and
9783    /// reshapes the caches of every layer whose KV head count differs from
9784    /// `num_kv_heads`. The loader and the tests share this one path, so a
9785    /// hand-built pipeline cannot hold a geometry the loader would refuse.
9786    /// Call before the first forward (it drops cached rows of reshaped
9787    /// layers). Refuses combinations whose paths would read it wrong:
9788    /// Gemma-4 global layers and MLA carry their own geometry.
9789    pub fn set_attn_geometry(
9790        &mut self,
9791        kv_heads_per_layer: Option<Vec<usize>>,
9792        v_head_dim: Option<usize>,
9793    ) -> Result<(), String> {
9794        if kv_heads_per_layer.is_some() || v_head_dim.is_some() {
9795            if self.global_attn.is_some() {
9796                return Err(
9797                    "per-layer KV heads / v_head_dim cannot combine with Gemma-4 global \
9798                     attention geometry"
9799                        .into(),
9800                );
9801            }
9802            if self
9803                .weights
9804                .layers
9805                .iter()
9806                .any(|lw| matches!(lw.attn, AttnKind::Mla(_)))
9807            {
9808                return Err("per-layer KV heads / v_head_dim cannot combine with MLA".into());
9809            }
9810        }
9811        if let Some(vd) = v_head_dim {
9812            if vd == 0 || vd > self.head_dim {
9813                return Err(format!(
9814                    "v_head_dim {vd} must be in 1..={} (head_dim)",
9815                    self.head_dim
9816                ));
9817            }
9818        }
9819        if let Some(v) = &kv_heads_per_layer {
9820            if v.len() != self.physical_layers {
9821                return Err(format!(
9822                    "kv_heads_per_layer has {} entries, expected {} layers",
9823                    v.len(),
9824                    self.physical_layers
9825                ));
9826            }
9827            for (li, &nkv) in v.iter().enumerate() {
9828                let is_attn = matches!(
9829                    self.weights.layers.get(li).map(|lw| &lw.attn),
9830                    Some(AttnKind::Full { .. }) | None
9831                );
9832                if !is_attn {
9833                    continue;
9834                }
9835                let nh = self
9836                    .attention_heads_per_layer
9837                    .as_ref()
9838                    .and_then(|h| h.get(li).copied())
9839                    .unwrap_or(self.num_heads);
9840                if nkv == 0 || nh % nkv != 0 {
9841                    return Err(format!(
9842                        "layer {li}: {nkv} KV heads must be nonzero and divide {nh} Q heads"
9843                    ));
9844                }
9845            }
9846        }
9847        self.kv_heads_per_layer = kv_heads_per_layer;
9848        self.v_head_dim = v_head_dim.filter(|&vd| vd != self.head_dim);
9849        if self.kv_heads_per_layer.is_some() {
9850            for li in 0..self.kv_cache.layers.len() {
9851                let full = matches!(
9852                    self.weights
9853                        .layers
9854                        .get(self.phys_layer(li))
9855                        .map(|lw| &lw.attn),
9856                    Some(AttnKind::Full { .. })
9857                );
9858                let nkv = self.layer_num_kv_heads(li);
9859                let cache = &self.kv_cache.layers[li];
9860                if full && (cache.num_kv_heads != nkv || cache.head_dim != self.head_dim) {
9861                    let sinks = cache.sinks.clone();
9862                    self.kv_cache.layers[li] =
9863                        crate::kv_cache::LayerKvCache::new(nkv, self.head_dim);
9864                    self.kv_cache.layers[li].sinks = sinks;
9865                }
9866            }
9867        }
9868        Ok(())
9869    }
9870
9871    /// Attach learned attention-sink logits (one per Q head) to PHYSICAL
9872    /// layer `phys` — every virtual layer that runs it. The loader calls
9873    /// this for each `model.layers.N.self_attn.sinks` tensor.
9874    pub fn set_layer_sinks(&mut self, phys: usize, sinks: Vec<f32>) -> Result<(), String> {
9875        let Some(lw) = self.weights.layers.get(phys) else {
9876            return Err(format!("sinks for layer {phys}: no such layer"));
9877        };
9878        if !matches!(lw.attn, AttnKind::Full { .. }) {
9879            return Err(format!(
9880                "sinks for layer {phys}: only softmax (Full) attention layers take sinks"
9881            ));
9882        }
9883        let nh = self
9884            .attention_heads_per_layer
9885            .as_ref()
9886            .and_then(|h| h.get(phys).copied())
9887            .unwrap_or(self.num_heads);
9888        if sinks.len() != nh {
9889            return Err(format!(
9890                "sinks for layer {phys}: {} values, expected one per Q head ({nh})",
9891                sinks.len()
9892            ));
9893        }
9894        if let Some(bad) = sinks.iter().find(|v| !v.is_finite()) {
9895            return Err(format!("sinks for layer {phys}: non-finite value {bad}"));
9896        }
9897        for li in 0..self.kv_cache.layers.len() {
9898            if self.phys_layer(li) == phys {
9899                self.kv_cache.layers[li].sinks = Some(sinks.clone());
9900            }
9901        }
9902        Ok(())
9903    }
9904
9905    /// Why the GPU attention graphs cannot serve this model, if they
9906    /// cannot: the wgpu whole-token and batched graphs, the greedy
9907    /// multi-burst, the q1 attention dropin and the Metal block/chunk/rows
9908    /// graphs all assume ONE (num_kv_heads, head_dim) geometry, V heads as
9909    /// wide as K, a single RoPE table, full-context attention and a plain
9910    /// softmax. A model outside that contract runs on the CPU layer walk
9911    /// (and the per-op GPU matvecs) until a graph learns it — never on a
9912    /// graph that would read it wrong. None = no attention-level reason
9913    /// (the graph builders still check weights and layer kinds).
9914    pub fn graph_attn_decline_reason(&self) -> Option<&'static str> {
9915        if self.kv_heads_per_layer.is_some() {
9916            return Some("per-layer KV head counts");
9917        }
9918        if self.v_head_dim.is_some_and(|vd| vd != self.head_dim) {
9919            return Some("V heads narrower than Q/K heads");
9920        }
9921        if self.kv_cache.layers.iter().any(|l| l.sinks.is_some()) {
9922            return Some("learned attention sinks");
9923        }
9924        if self.swa.is_some() || self.sliding_layers.is_some() {
9925            return Some("sliding-window layers");
9926        }
9927        None
9928    }
9929
9930    /// Why the WGPU graphs (whole-token, batched prefill, greedy burst)
9931    /// cannot run this model's attention, if they cannot. Per-layer KV
9932    /// heads, V narrower than K, learned sinks and sliding windows ride
9933    /// their per-layer geometry (`GraphAttnGeom`, the ATTEND_X kernels);
9934    /// what that geometry does not express keeps the decline, by name.
9935    /// None for every model with one attention geometry.
9936    pub fn wgpu_graph_attn_decline(&self) -> Option<&'static str> {
9937        self.graph_attn_decline_reason()?;
9938        if self.global_attn.is_some() {
9939            return Some("per-layer head width (Gemma-4 global layers) with per-layer geometry");
9940        }
9941        if self.attention_heads_per_layer.is_some() {
9942            return Some("per-layer Q head counts with per-layer geometry");
9943        }
9944        if self.attn_v_norm {
9945            return Some("V norm with per-layer geometry");
9946        }
9947        if (0..self.num_layers).any(|li| self.layer_rope_scale(li) != 1.0) {
9948            return Some("scaled RoPE positions with per-layer geometry");
9949        }
9950        if self.weights.layers.iter().any(|lw| {
9951            lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some()
9952        }) {
9953            return Some("sandwich norms / layer scale with per-layer geometry");
9954        }
9955        if self.weights.layers.iter().any(|lw| {
9956            matches!(
9957                &lw.attn,
9958                AttnKind::Full {
9959                    output_gate: true,
9960                    ..
9961                }
9962            )
9963        }) && self.v_head_dim.is_some()
9964        {
9965            return Some("gated attention with V narrower than K");
9966        }
9967        if (0..self.num_layers).any(|li| {
9968            self.layer_is_local(li)
9969                && self.inv_freq_local.is_none()
9970                && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
9971        }) {
9972            return Some("local rotary width without a local RoPE table");
9973        }
9974        None
9975    }
9976
9977    /// The wgpu graphs' attention geometry for layer `li` (virtual index):
9978    /// Some only for a model whose layers do not share one (MiMo-V2) — KV
9979    /// heads, V width, rotary width and RoPE table, window and sinks of
9980    /// THIS layer, exactly what the CPU attention reads for it.
9981    fn graph_attn_geom(&self, li: usize) -> Option<crate::gpu::GraphAttnGeom<'_>> {
9982        self.graph_attn_decline_reason()?;
9983        let (nkv, _hd, rd) = self.layer_geom(li);
9984        let invf: &[f32] = if self.layer_is_local(li) {
9985            match &self.inv_freq_local {
9986                Some(f) => f.as_slice(),
9987                None => self.inv_freq.as_slice(),
9988            }
9989        } else {
9990            match &self.inv_freq_global {
9991                Some(f) => f.as_slice(),
9992                None => self.inv_freq.as_slice(),
9993            }
9994        };
9995        Some(crate::gpu::GraphAttnGeom {
9996            nkv,
9997            dv: self.layer_v_dim(li),
9998            rd,
9999            invf,
10000            window: self.layer_window(li),
10001            sink: self.kv_cache.layers[li].sinks.as_deref(),
10002        })
10003    }
10004
10005    /// Bring the host KV cache of every Full-attention layer in
10006    /// `[from, upto)` up to `position` rows from the wgpu mirrors, where a
10007    /// device graph advanced a layer that the host is about to run: a
10008    /// device prefix that shrank since the prompt (or a batched prefill
10009    /// prefix longer than the decode one). A layer whose mirror does not
10010    /// hold the missing rows is left alone. Rows a sliding layer's ring
10011    /// no longer holds come back as zeros — outside every window that
10012    /// will read them.
10013    #[cfg(feature = "gpu")]
10014    fn pull_lagging_host_kv(&mut self, from: usize, upto: usize, position: usize) {
10015        let kv_id = self.graph_kv_id;
10016        for li in from..upto.min(self.num_layers) {
10017            if !matches!(
10018                self.weights.layers[self.phys_layer(li)].attn,
10019                AttnKind::Full { .. }
10020            ) {
10021                continue;
10022            }
10023            let host = self.kv_cache.layers[li].seq_len;
10024            if host >= position {
10025                continue;
10026            }
10027            let Some(dev) = crate::gpu::graph_kv_stored(kv_id, li) else {
10028                continue;
10029            };
10030            let to = dev.min(position);
10031            if to <= host {
10032                continue;
10033            }
10034            let (nkv, hd) = {
10035                let c = &self.kv_cache.layers[li];
10036                (c.num_kv_heads, c.head_dim)
10037            };
10038            let Some((k, v, first_valid)) =
10039                crate::gpu::graph_kv_pull_host(kv_id, li, host, to, nkv, hd)
10040            else {
10041                continue;
10042            };
10043            // A sliding layer only ever reads its last `window` rows; a
10044            // full-context layer needs every row it did not have.
10045            let need_from = match self.layer_window(li) {
10046                Some(w) => host.max((position + 1).saturating_sub(w)),
10047                None => host,
10048            };
10049            if first_valid > need_from {
10050                tracing::warn!(
10051                    "layer {li}: device KV rows {host}..{to} no longer resident \
10052                     (from {first_valid}); host attention will miss them"
10053                );
10054            }
10055            let row = nkv * hd;
10056            let cache = &mut self.kv_cache.layers[li];
10057            for p in 0..to - host {
10058                cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
10059            }
10060        }
10061    }
10062
10063    /// Log (once per graph site and pipeline) that `site` declined for
10064    /// `reason`. The lines are kept so a caller or a test can read them.
10065    fn note_graph_decline(&self, site: &'static str, reason: &'static str) {
10066        let mut seen = self.graph_declines.borrow_mut();
10067        if !seen.iter().any(|&(s, r)| s == site && r == reason) {
10068            tracing::warn!("{site} declined: {reason} (CPU attention path)");
10069            seen.push((site, reason));
10070        }
10071    }
10072
10073    /// The GPU-graph declines this pipeline has logged so far, as the
10074    /// logged lines.
10075    pub fn graph_declines(&self) -> Vec<String> {
10076        self.graph_declines
10077            .borrow()
10078            .iter()
10079            .map(|(site, reason)| format!("{site} declined: {reason} (CPU attention path)"))
10080            .collect()
10081    }
10082
10083    /// Does layer `li` have the plain attention geometry the historical
10084    /// head-masked f32 path (`multi_head_attention`) assumes — pipeline-wide
10085    /// KV heads / head_dim / RoPE table, full context, no sink, V as wide
10086    /// as K? Anything else runs the dense `qwen_attention` instead.
10087    /// `CMF_LAYER_DUMP` writer (see `Pipeline::layer_dump`): one position's
10088    /// hidden after layer `li` as raw little-endian f32 into
10089    /// `<dir>/p{pos:06}_l{li:02}.f32`. A failed write is reported once and
10090    /// never stops the forward — the dump is a diagnostic.
10091    fn dump_layer_row(&self, pos: usize, li: usize, row: &[f32]) {
10092        let Some(dir) = &self.layer_dump else {
10093            return;
10094        };
10095        let mut bytes = Vec::with_capacity(row.len() * 4);
10096        for v in row {
10097            bytes.extend_from_slice(&v.to_le_bytes());
10098        }
10099        let path = dir.join(format!("p{pos:06}_l{li:02}.f32"));
10100        if let Err(e) = std::fs::create_dir_all(dir).and_then(|_| std::fs::write(&path, &bytes)) {
10101            use std::sync::atomic::{AtomicBool, Ordering};
10102            static SAID: AtomicBool = AtomicBool::new(false);
10103            if !SAID.swap(true, Ordering::Relaxed) {
10104                tracing::error!("CMF_LAYER_DUMP: cannot write {}: {e}", path.display());
10105            }
10106        }
10107    }
10108
10109    /// Decide the MiMo-V2 expert placement once (`crate::mimo_moe`). Any
10110    /// other model turns the slot off on the first call.
10111    fn mimo_moe_prepare(&mut self) {
10112        if !self.mimo_moe.is_undecided() {
10113            return;
10114        }
10115        let slot = {
10116            let layers: Vec<(usize, &MoeFfn)> = (0..self.num_layers)
10117                .filter_map(
10118                    |li| match &self.weights.layers.get(self.phys_layer(li))?.ffn {
10119                        FfnKind::Moe(m) => Some((li, m)),
10120                        _ => None,
10121                    },
10122                )
10123                .collect();
10124            // One bank lives on one device: an in-process multi-GPU split
10125            // keeps the whole-layer path.
10126            if layers.is_empty()
10127                || self.physical_layers != self.num_layers
10128                || self.gpu_plan.is_some()
10129            {
10130                crate::mimo_moe::Slot::Off
10131            } else {
10132                // Whether a whole-token graph could run this model's layers
10133                // (then a whole-layer prefix is one submit, not per-layer
10134                // fences).
10135                let graph_prefix =
10136                    self.wgpu_graph_attn_decline().is_none() && crate::gpu::wgpu_graph_default();
10137                crate::mimo_moe::Slot::decide(&layers, self.num_layers, graph_prefix)
10138            }
10139        };
10140        self.mimo_moe = slot;
10141    }
10142
10143    #[cfg(test)]
10144    pub(crate) fn test_graph_kv_id(&self) -> u64 {
10145        self.graph_kv_id
10146    }
10147
10148    /// Dynamic MiMo layer: one device attention graph, followed by a bank
10149    /// frame. Both decode and short verification use this same attention
10150    /// path and absolute layer key; the host KV may intentionally lag.
10151    pub(crate) fn mimo_graph_layer_rows(
10152        &mut self,
10153        li: usize,
10154        h: &mut [f32],
10155        positions: &[usize],
10156    ) -> crate::gpu::BatchGraphOutcome {
10157        use crate::gpu::BatchGraphOutcome as Out;
10158        let b = positions.len();
10159        if !(1..=4).contains(&b)
10160            || h.len() != b * self.hidden_size
10161            || !self.mimo_moe.is_dynamic(li, true)
10162            || !crate::gpu::enabled_here()
10163            || !crate::gpu::wgpu_active()
10164            || self.o1_active()
10165            || self.physical_layers != self.num_layers
10166            // The pair-fusion diagnostic (and an explicit graph-off run)
10167            // rewinds only host KV. A hidden singleton attention graph here
10168            // would leave device mirrors ahead of the next host position.
10169            || std::env::var("CMF_GPU_WGPU_GRAPH").as_deref() == Ok("0")
10170            || std::env::var("CMF_MIMO_ATTN_GRAPH").as_deref() == Ok("0")
10171            || self.wgpu_graph_attn_decline().is_some()
10172        {
10173            return Out::Declined;
10174        }
10175        let attn_started = std::time::Instant::now();
10176        let outcome = {
10177            let lw = &self.weights.layers[li];
10178            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
10179                return Out::Declined;
10180            }
10181            let FfnKind::Moe(m) = &lw.ffn else {
10182                return Out::Declined;
10183            };
10184            let AttnKind::Full {
10185                wq,
10186                wk,
10187                wv,
10188                wo,
10189                q_norm,
10190                k_norm,
10191                output_gate,
10192                softplus_gate,
10193                bias,
10194            } = &lw.attn
10195            else {
10196                return Out::Declined;
10197            };
10198            if *output_gate || softplus_gate.is_some() {
10199                return Out::Declined;
10200            }
10201            let Some(model) = wq.graph_weight().map(|(m, ..)| m.clone()).or_else(|| {
10202                m.experts
10203                    .first()?
10204                    .gate_proj
10205                    .mapped_q4tp()
10206                    .map(|(m, _)| m.clone())
10207            }) else {
10208                return Out::Declined;
10209            };
10210            fn gw<'a>(
10211                t: &'a QTensor,
10212                owner: &std::sync::Arc<cortiq_core::CmfModel>,
10213            ) -> Option<crate::gpu::GraphW<'a>> {
10214                if let Some((m, idx, kind, rs)) = t.graph_weight() {
10215                    if m.uid() != owner.uid() || t.has_prism_contract() {
10216                        return None;
10217                    }
10218                    return Some(crate::gpu::GraphW {
10219                        idx,
10220                        kind,
10221                        row_scale: rs,
10222                        data: &[],
10223                        prism: crate::gpu::GraphPrismOp::None,
10224                        affine: false,
10225                    });
10226                }
10227                t.as_f32().map(|data| crate::gpu::GraphW {
10228                    idx: 0,
10229                    kind: 4,
10230                    row_scale: &[],
10231                    data,
10232                    prism: crate::gpu::GraphPrismOp::None,
10233                    affine: false,
10234                })
10235            }
10236            let (Some(q), Some(k), Some(v), Some(o)) = (
10237                gw(wq, &model),
10238                gw(wk, &model),
10239                gw(wv, &model),
10240                gw(wo, &model),
10241            ) else {
10242                return Out::Declined;
10243            };
10244            let layer = crate::gpu::GraphLayer {
10245                input_norm: &lw.input_norm,
10246                post_norm: &lw.post_norm,
10247                ffn: crate::gpu::GraphFfn::AttentionOnly,
10248                attn: crate::gpu::GraphAttn::Full {
10249                    wq: q,
10250                    wk: k,
10251                    wv: v,
10252                    wo: o,
10253                    q_norm: q_norm.as_deref(),
10254                    k_norm: k_norm.as_deref(),
10255                    late_qk_norm: self.qk_norm_after_rope,
10256                    bias: bias
10257                        .as_ref()
10258                        .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice())),
10259                    output_gate: false,
10260                    cpu_k: self.kv_cache.layers[li].k_heads(),
10261                    cpu_v: self.kv_cache.layers[li].v_heads(),
10262                    geom: self.graph_attn_geom(li),
10263                },
10264            };
10265            let (nkv, hd, rd) = self.layer_geom(li);
10266            crate::gpu::forward_batch_graph_at(
10267                &model,
10268                self.graph_kv_id,
10269                li,
10270                &[layer],
10271                &self.inv_freq,
10272                h,
10273                self.layer_num_heads(li),
10274                nkv,
10275                hd,
10276                rd,
10277                self.hidden_size,
10278                1,
10279                positions,
10280                self.kv_cache.max_seq_len,
10281                self.norm_style == cortiq_core::NormStyle::Gemma,
10282                self.rms_eps as f32,
10283                self.attn_scale,
10284                b,
10285                &[],
10286                self.o1_epoch,
10287                None,
10288                None,
10289            )
10290        };
10291        match outcome {
10292            Out::Completed => {}
10293            Out::Declined => return Out::Declined,
10294            Out::Failed => {
10295                self.graph_failed
10296                    .store(true, std::sync::atomic::Ordering::Relaxed);
10297                return Out::Failed;
10298            }
10299        }
10300        let attn_ns = attn_started.elapsed().as_nanos() as u64;
10301        let hs = self.hidden_size;
10302        let lw = &self.weights.layers[li];
10303        let FfnKind::Moe(m) = &lw.ffn else {
10304            unreachable!()
10305        };
10306        let mut post = vec![0.0; h.len()];
10307        for (x, y) in h.chunks_exact(hs).zip(post.chunks_exact_mut(hs)) {
10308            inference::rms_norm_into(x, &lw.post_norm, self.rms_eps, self.norm_style, y);
10309        }
10310        let mut ffn = if b == 1 {
10311            moe_ffn_banked(&mut self.mimo_moe, li, m, &post, self.pool.as_deref())
10312        } else {
10313            moe_ffn_banked_rows(
10314                &mut self.mimo_moe,
10315                li,
10316                m,
10317                &post,
10318                b,
10319                hs,
10320                self.pool.as_deref(),
10321            )
10322        };
10323        for (x, &f) in h.iter_mut().zip(&ffn) {
10324            *x += f;
10325        }
10326        attention::recycle_buf(&mut ffn);
10327        if self.layer_dump.is_some() {
10328            for (&pos, row) in positions.iter().zip(h.chunks_exact(hs)) {
10329                self.dump_layer_row(pos, li, row);
10330            }
10331        }
10332        crate::mimo_moe::note_attention_graph(b, attn_ns);
10333        Out::Completed
10334    }
10335
10336    fn layer_attn_plain(&self, li: usize) -> bool {
10337        self.kv_heads_per_layer.is_none()
10338            && self.v_head_dim.is_none()
10339            && self.global_attn.is_none()
10340            && self.layer_window(li).is_none()
10341            && self.kv_cache.layers[li].sinks.is_none()
10342    }
10343
10344    /// Forward one position through all layers (hybrid dispatch).
10345    fn forward_layers(
10346        &mut self,
10347        hidden: &[f32],
10348        position: usize,
10349        task_mask: Option<&TaskMask>,
10350    ) -> Vec<f32> {
10351        let out = self.forward_layers_upto(hidden, position, task_mask, None);
10352        self.o1_progress();
10353        out
10354    }
10355
10356    // ── Network pipeline-split building blocks (coordinator/worker) ──
10357    // A remote worker owns layers [from ..= upto] and their KV; the
10358    // coordinator owns the rest plus embed / final norm / head. Attention
10359    // causality is per-layer, so a whole prompt's boundary hiddens ship
10360    // as one batch and decode ships one vector per token.
10361
10362    /// Embed one token id (embed multiplier applied).
10363    pub fn embed_id(&self, id: u32) -> Vec<f32> {
10364        self.embed_single(id)
10365    }
10366
10367    /// Refuse the archs/modes whose forward cannot be cut at a layer
10368    /// boundary. Loud by design: a split that silently changed the math
10369    /// would be a chimera.
10370    pub fn split_supported(&self) -> Result<(), String> {
10371        if self.dsv4.is_some() {
10372            return Err(
10373                "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
10374            );
10375        }
10376        if self.dsv41.is_some() {
10377            return Err(
10378                "network split: DeepSeek-V4.1 owns the shared CED/CSA2 state (not splittable)"
10379                    .into(),
10380            );
10381        }
10382        if self.qwen4_exp.is_some() {
10383            return Err(
10384                "network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
10385            );
10386        }
10387        if self.g3n.is_some() {
10388            return Err(
10389                "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
10390            );
10391        }
10392        Ok(())
10393    }
10394
10395    /// Forward `hidden` through layers [from ..= upto] at `position`,
10396    /// appending those layers' KV/state. Both split sides call this
10397    /// over their own range; a task mask applies to the span's own
10398    /// layers (each side masks what it runs).
10399    pub fn forward_span(
10400        &mut self,
10401        hidden: &[f32],
10402        position: usize,
10403        from: usize,
10404        upto: usize,
10405        task_mask: Option<&TaskMask>,
10406    ) -> Result<Vec<f32>, String> {
10407        self.split_supported()?;
10408        if from > upto || upto >= self.num_layers {
10409            return Err(format!(
10410                "forward_span: layer range {from}..={upto} outside 0..{}",
10411                self.num_layers
10412            ));
10413        }
10414        if hidden.len() != self.hidden_size {
10415            return Err(format!(
10416                "forward_span: hidden len {} ≠ hidden_size {}",
10417                hidden.len(),
10418                self.hidden_size
10419            ));
10420        }
10421        let out = self.forward_layers_span(hidden, position, task_mask, from, Some(upto));
10422        self.o1_progress();
10423        if self
10424            .graph_failed
10425            .swap(false, std::sync::atomic::Ordering::Relaxed)
10426        {
10427            self.cancel
10428                .store(false, std::sync::atomic::Ordering::Relaxed);
10429            self.clear_sequence_state();
10430            return Err("forward_span: deferred O(1) transition failed".into());
10431        }
10432        Ok(out)
10433    }
10434
10435    /// Final norm + lm_head over a boundary hidden (the final-logit
10436    /// softcap is applied by lm_head_forward itself).
10437    pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
10438        let normed = inference::rms_norm(
10439            hidden,
10440            &self.weights.final_norm,
10441            self.rms_eps,
10442            self.norm_style,
10443        );
10444        self.lm_head_forward(&normed)
10445    }
10446
10447    /// Sample the next token with this pipeline's sampler state.
10448    pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
10449        sampler::sample_with_scratch(
10450            logits,
10451            &self.sampler_config,
10452            past_tokens,
10453            &mut self.rng,
10454            &mut self.sampler_scratch,
10455        )
10456    }
10457
10458    /// Fresh sequence: clear KV, reuse history and device mirrors.
10459    pub fn reset_session(&mut self) {
10460        self.clear_sequence_state();
10461    }
10462
10463    /// Batched span prefill from token ids (coordinator side): embed +
10464    /// layers [0 ..= upto]; returns the boundary hiddens of ALL positions
10465    /// (ids.len() × hidden). Rides the same layer-major machinery as the
10466    /// local prefill; falls back to the per-position walk under
10467    /// CMF_PREFILL=seq.
10468    pub fn prefill_span_ids(
10469        &mut self,
10470        ids: &[u32],
10471        start_pos: usize,
10472        upto: usize,
10473        task_mask: Option<&TaskMask>,
10474    ) -> Result<Vec<f32>, String> {
10475        self.split_supported()?;
10476        if upto >= self.num_layers {
10477            return Err(format!(
10478                "prefill_span_ids: upto {upto} outside 0..{}",
10479                self.num_layers
10480            ));
10481        }
10482        // Same predicate as the whole-stack prefill: a span whose GDN
10483        // state lives on the device must walk positions through the
10484        // graph, not through the batched CPU span.
10485        if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10486            let out =
10487                self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1);
10488            self.check_o1_progress_failure("prefill_span_ids")?;
10489            Ok(out)
10490        } else {
10491            let hs = self.hidden_size;
10492            let mut out = Vec::with_capacity(ids.len() * hs);
10493            for (i, &id) in ids.iter().enumerate() {
10494                let emb = self.embed_id(id);
10495                out.extend_from_slice(&self.forward_span(
10496                    &emb,
10497                    start_pos + i,
10498                    0,
10499                    upto,
10500                    task_mask,
10501                )?);
10502            }
10503            Ok(out)
10504        }
10505    }
10506
10507    /// Batched span prefill from boundary hiddens (worker side): layers
10508    /// [from ..= upto] for every position in the batch; returns the batch.
10509    pub fn prefill_span_hidden(
10510        &mut self,
10511        hidden: &[f32],
10512        start_pos: usize,
10513        from: usize,
10514        upto: usize,
10515        task_mask: Option<&TaskMask>,
10516    ) -> Result<Vec<f32>, String> {
10517        self.split_supported()?;
10518        let hs = self.hidden_size;
10519        if hidden.is_empty() || hidden.len() % hs != 0 {
10520            return Err(format!(
10521                "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
10522                hidden.len()
10523            ));
10524        }
10525        if from > upto || upto >= self.num_layers {
10526            return Err(format!(
10527                "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
10528                self.num_layers
10529            ));
10530        }
10531        if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10532            let out = self.prefill_batch_span(
10533                PrefillIn::Hidden(hidden),
10534                start_pos,
10535                task_mask,
10536                from,
10537                upto + 1,
10538            );
10539            self.check_o1_progress_failure("prefill_span_hidden")?;
10540            Ok(out)
10541        } else {
10542            let b = hidden.len() / hs;
10543            let mut out = Vec::with_capacity(hidden.len());
10544            for i in 0..b {
10545                let h = self.forward_span(
10546                    &hidden[i * hs..(i + 1) * hs],
10547                    start_pos + i,
10548                    from,
10549                    upto,
10550                    task_mask,
10551                )?;
10552                out.extend_from_slice(&h);
10553            }
10554            Ok(out)
10555        }
10556    }
10557
10558    /// Build the whole-token wgpu graph for a pure-attention q1 model (every
10559    /// layer Full q1 + dense q1 FFN, no gate/bias). Returns the post-stack
10560    /// hidden (caller does final norm + lm_head), or None to fall back.
10561    fn try_token_graph_wgpu(
10562        &self,
10563        hidden: &[f32],
10564        position: usize,
10565        logits_out: &mut Vec<f32>,
10566        layers_run: &mut usize,
10567    ) -> Option<Result<Vec<f32>, ()>> {
10568        self.try_token_graph_wgpu_steps(
10569            hidden,
10570            position,
10571            logits_out,
10572            1,
10573            None,
10574            Some(layers_run),
10575            0,
10576            self.num_layers,
10577        )
10578    }
10579
10580    /// The span twin (network split): the graph covers [from..upto_excl)
10581    /// — one submit per SEGMENT per token. lm_head folds in only when
10582    /// the span reaches the last layer.
10583    fn try_token_graph_wgpu_span(
10584        &self,
10585        hidden: &[f32],
10586        position: usize,
10587        logits_out: &mut Vec<f32>,
10588        from: usize,
10589        upto_excl: usize,
10590        layers_run: &mut usize,
10591    ) -> Option<Result<Vec<f32>, ()>> {
10592        self.try_token_graph_wgpu_steps(
10593            hidden,
10594            position,
10595            logits_out,
10596            1,
10597            None,
10598            Some(layers_run),
10599            from,
10600            upto_excl,
10601        )
10602    }
10603
10604    /// Greedy burst: forward `t_next` and let the device pick + re-embed
10605    /// the next k−1 tokens — k frames, ONE submit, k ids back. The ZML
10606    /// trade, on wgpu. None ⇒ caller keeps the per-token path.
10607    fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
10608        if self.o1_active() || self.attn_softcap > 0.0 {
10609            return None;
10610        }
10611        // The burst builds the whole-token graph; attention the graph's
10612        // per-layer geometry cannot express keeps the per-token path.
10613        if let Some(reason) = self.wgpu_graph_attn_decline() {
10614            self.note_graph_decline("wgpu multi-burst", reason);
10615            return None;
10616        }
10617        let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
10618        if !graph_on || self.graph_refused() {
10619            // Same memo as the decode site: this path builds the very
10620            // same graph, so a model it cannot build for must not be
10621            // walked again here either. Missing this guard was worth
10622            // 2.5x on an Adreno — 0.361 tok/s against 0.905 — because
10623            // the burst retried per token what decode had already given
10624            // up on.
10625            return None;
10626        }
10627        let emb = self.embed_single(t_next);
10628        let mut lg = Vec::new();
10629        let mut ids = Vec::new();
10630        match self.try_token_graph_wgpu_steps(
10631            &emb,
10632            position,
10633            &mut lg,
10634            k,
10635            Some(&mut ids),
10636            None,
10637            0,
10638            self.num_layers,
10639        ) {
10640            Some(Ok(_)) => {}
10641            Some(Err(())) => {
10642                // Preserve the backend's post-admission failure through the
10643                // Option-based burst API.  The decode caller consumes this
10644                // flag and clears the sequence instead of falling through
10645                // to a stale CPU recurrent state.
10646                self.graph_failed
10647                    .store(true, std::sync::atomic::Ordering::Relaxed);
10648                return None;
10649            }
10650            None => return None,
10651        }
10652        (ids.len() == k).then_some(ids)
10653    }
10654
10655    /// Multi-step greedy: k whole frames in ONE submit, argmax and re-embed
10656    /// on the device. `ids_out` receives the k winner ids; the hidden/logits
10657    /// outputs are NOT produced in that mode.
10658    fn try_token_graph_wgpu_steps(
10659        &self,
10660        hidden: &[f32],
10661        position: usize,
10662        logits_out: &mut Vec<f32>,
10663        steps: usize,
10664        ids_out: Option<&mut Vec<u32>>,
10665        layers_run: Option<&mut usize>,
10666        from: usize,
10667        upto_excl: usize,
10668    ) -> Option<Result<Vec<f32>, ()>> {
10669        // The bank has already reserved its VRAM. Never build a second
10670        // all-expert arena across bank-owned layers (including bursts).
10671        let upto_excl = match self.mimo_moe.graph_prefix_end() {
10672            Some(end) if end < upto_excl => {
10673                if steps != 1 || layers_run.is_none() || from >= end {
10674                    return None;
10675                }
10676                end
10677            }
10678            _ => upto_excl,
10679        };
10680        // O(1) Nyström decode runs off the sealed state, not the KV cache the
10681        // graph mirrors — never take the graph while o1 is active.
10682        let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
10683        if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
10684            // Softcapped scores have no graph kernel yet — CPU owns them.
10685            // o1 rides the graph only behind CMF_O1_GPU=1 while the port
10686            // proves itself; without it the CPU path owns o1 as before.
10687            return None;
10688        }
10689        // Per-layer KV heads, narrow V, sinks and sliding windows ride
10690        // `GraphAttn::Full::geom` (the ATTEND_X kernels). Anything that
10691        // geometry cannot express declines here, by name — before the
10692        // per-layer gate existed a sliding/sink model ran the graph as
10693        // full-context attention, fluent and wrong. The caller memoizes
10694        // the refusal.
10695        if let Some(reason) = self.wgpu_graph_attn_decline() {
10696            self.note_graph_decline("wgpu token graph", reason);
10697            return None;
10698        }
10699        // Per-layer sealed o1 state for the graph. During prefill the
10700        // state is still Collecting -> views are None -> the graph
10701        // refuses below and the CPU prefill records the q trace and
10702        // seals, exactly as the o1 design requires.
10703        let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
10704            .map(|li| {
10705                if !o1_gpu {
10706                    return None;
10707                }
10708                self.kv_cache.layers[self.phys_layer(li)].o1_views()
10709            })
10710            .collect();
10711        if self.o1_active() && o1_gpu {
10712            // Any o1 layer not sealed (or degenerate exact-only) keeps the
10713            // whole token on the CPU: half-graph forwards would desync.
10714            let want: usize = (from..upto_excl)
10715                .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
10716                .count();
10717            let have = o1_views.iter().filter(|v| v.is_some()).count();
10718            if want == 0 || have != want {
10719                // The silent twin of the gpu-side o1 gates, found the
10720                // same way: a 15x decode drop with an empty log. Views
10721                // stay None until the layer's state SEALS, so `have`
10722                // lagging `want` early in a run is the o1 design working
10723                // — but it must say so, or the next reader spends a
10724                // night proving the kernels innocent.
10725                // On CHANGE, not once: the first decline is the legal
10726                // unsealed prefill, and a once-print buries the state
10727                // that matters — what the count reads AFTER the seal.
10728                use std::sync::atomic::{AtomicUsize, Ordering};
10729                static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
10730                let code = have * 1000 + want;
10731                if LAST.swap(code, Ordering::Relaxed) != code {
10732                    tracing::warn!(
10733                        "o1 graph: {have} of {want} layers sealed — per-op until all seal"
10734                    );
10735                }
10736                return None;
10737            }
10738        }
10739        let nh = self.num_heads;
10740        let (nkv, hd, rd) = self.layer_geom(0);
10741        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
10742        let mut layers = Vec::with_capacity(upto_excl - from);
10743        let mut model = None;
10744        let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
10745        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
10746            if let Some((m, i, kind, rs)) = t
10747                .graph_weight()
10748                .or_else(|| t.graph_weight_descriptor())
10749            {
10750                let name = &m.tensors[i].name;
10751                let prism = if crate::prism::is_inverse_embedding(m, name) {
10752                    crate::gpu::GraphPrismOp::InverseEmbedding
10753                } else if crate::prism::is_forward_weight(m, name) {
10754                    crate::gpu::GraphPrismOp::Forward
10755                } else {
10756                    crate::gpu::GraphPrismOp::None
10757                };
10758                return Some(crate::gpu::GraphW {
10759                    idx: i,
10760                    kind,
10761                    row_scale: rs,
10762                    data: &[],
10763                    prism,
10764                    affine: crate::prism::is_affine_target(m, name),
10765                });
10766            }
10767            // Small unquantized projections (GDN in_proj_a/b) stay f32.
10768            match t.as_f32() {
10769                Some(d) => Some(crate::gpu::GraphW {
10770                    idx: 0,
10771                    kind: 4,
10772                    row_scale: &[],
10773                    data: d,
10774                    prism: crate::gpu::GraphPrismOp::None,
10775                    affine: false,
10776                }),
10777                None => {
10778                    if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
10779                        eprintln!("batch graph: weight has no graph/f32 representation");
10780                    }
10781                    None
10782                }
10783            }
10784        }
10785        for li in from..upto_excl {
10786            let lw = &self.weights.layers[self.phys_layer(li)];
10787            if dbg {
10788                let ak = match &lw.attn {
10789                    AttnKind::Mla(_) => "Mla".into(),
10790                    AttnKind::Full {
10791                        output_gate, bias, ..
10792                    } => format!("Full gate={output_gate} bias={}", bias.is_some()),
10793                    AttnKind::LinearGdn(_) => "LinearGdn".into(),
10794                    AttnKind::Kda(_) => "Kda".into(),
10795                    AttnKind::Linear(_) => "Linear".into(),
10796                    AttnKind::ShortConv(_) => "ShortConv".into(),
10797                    AttnKind::Bounded(_) => "Bounded".into(),
10798                };
10799                let fk = match &lw.ffn {
10800                    FfnKind::Dense(_) => "Dense",
10801                    FfnKind::Moe(_) => "Moe",
10802                    FfnKind::DenseMoe(_) => "DenseMoe",
10803                };
10804                eprintln!("graph L{li}: attn={ak} ffn={fk}");
10805            }
10806            let gffn = match &lw.ffn {
10807                FfnKind::DenseMoe(_) => return None, // dual branch: CPU path
10808                // A tube layer is several matrices, not one — the
10809                // whole-layer graph has no shape for it yet.
10810                FfnKind::Dense(d) if !d.segs.is_empty() => return None,
10811                FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
10812                    gate: gw(&d.gate_proj)?,
10813                    up: gw(&d.up_proj)?,
10814                    down: gw(&d.down_proj)?,
10815                },
10816                FfnKind::Moe(m) => {
10817                    // Adaptive τ and expert masks keep the CPU path, where
10818                    // they are implemented. Sigmoid routing with a selection
10819                    // bias (LFM2-MoE / DeepSeek noaux_tc), a routed scale ≠ 1
10820                    // and an UNGATED shared expert (HunYuan hy_v3: ×2.826 on
10821                    // the routed mix, the shared expert at weight 1) are all
10822                    // graphed — before, every such token fell to the per-op
10823                    // path whole (145 submits/token on Hy-MT2-30B-A3B).
10824                    if m.route_tau.is_some() || m.mask.is_some() {
10825                        return None;
10826                    }
10827                    let shared = m.shared.as_ref();
10828                    let has_shared = shared.is_some();
10829                    let shared_gated = matches!(shared, Some((_, Some(_))));
10830                    let sgate = match shared {
10831                        Some((_, Some(sg))) => gw(sg)?,
10832                        // No gate (hy_v3) or no shared expert at all: the
10833                        // router weight stands in so the plumbing stays
10834                        // total; the select kernels pin weight 1 or skip.
10835                        _ => gw(&m.router)?,
10836                    };
10837                    let router = gw(&m.router)?;
10838                    // The resident MoE kernels do not yet carry the
10839                    // descriptor-aware transform through router/shared-gate
10840                    // selection.  Refuse the complete layer instead of
10841                    // scoring with an untransformed Prism plane (the dense
10842                    // path has an explicit FWHT boundary below).
10843                    if router.prism != crate::gpu::GraphPrismOp::None
10844                        || sgate.prism != crate::gpu::GraphPrismOp::None
10845                        || router.affine
10846                        || sgate.affine
10847                    {
10848                        tracing::warn!(
10849                            "resident MoE declined: Prism/affine router or shared gate transform is not implemented"
10850                        );
10851                        return None;
10852                    }
10853                    let inter = m.experts.first()?.gate_proj.rows();
10854                    let mut experts = Vec::with_capacity(m.experts.len() + 1);
10855                    // q4t or q4tp, but not both in one layer — the kernels
10856                    // are picked per layer, not per expert.
10857                    let mut q4tp: Option<bool> = None;
10858                    // The mixed 2-bit profile: q2tp gate/up over a q4tp
10859                    // down. Uniform across the layer, like `q4tp` itself.
10860                    let mut gu_q2: Option<bool> = None;
10861                    for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
10862                        if !matches!(e.act, Act::Silu)
10863                            || e.gate_proj.rows() != inter
10864                            || e.up_proj.rows() != inter
10865                        {
10866                            return None;
10867                        }
10868                        // Expert tensors are packed into one resident buffer
10869                        // and the MoE kernels have no transform slot per
10870                        // expert.  Keep the CPU/per-op owner for Prism or
10871                        // affine experts rather than silently using raw bytes.
10872                        for expert_weight in [&e.gate_proj, &e.up_proj, &e.down_proj] {
10873                            let Some((em, ei, _, _)) = expert_weight
10874                                .graph_weight()
10875                                .or_else(|| expert_weight.graph_weight_descriptor())
10876                            else {
10877                                return None;
10878                            };
10879                            let name = &em.tensors[ei].name;
10880                            if crate::prism::is_forward_weight(em, name)
10881                                || crate::prism::is_inverse_embedding(em, name)
10882                                || crate::prism::is_affine_target(em, name)
10883                            {
10884                                tracing::warn!(
10885                                    "resident MoE declined: expert Prism/affine transform is not implemented"
10886                                );
10887                                return None;
10888                            }
10889                        }
10890                        let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
10891                            Some((mm, gi)) => (
10892                                mm,
10893                                gi,
10894                                e.up_proj.mapped_q4t()?.1,
10895                                e.down_proj.mapped_q4t()?.1,
10896                                false,
10897                                false,
10898                            ),
10899                            None => match e.gate_proj.mapped_q2tp() {
10900                                Some((mm, gi)) => (
10901                                    mm,
10902                                    gi,
10903                                    e.up_proj.mapped_q2tp()?.1,
10904                                    e.down_proj.mapped_q4tp()?.1,
10905                                    true,
10906                                    true,
10907                                ),
10908                                None => {
10909                                    let (mm, gi) = e.gate_proj.mapped_q4tp()?;
10910                                    (
10911                                        mm,
10912                                        gi,
10913                                        e.up_proj.mapped_q4tp()?.1,
10914                                        e.down_proj.mapped_q4tp()?.1,
10915                                        true,
10916                                        false,
10917                                    )
10918                                }
10919                            },
10920                        };
10921                        if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
10922                        {
10923                            // The shared expert rides in the same packed
10924                            // buffer as the routed ones, so a layer that
10925                            // mixes layouts cannot be indexed by one stride.
10926                            // Say so: the symptom is a whole model quietly
10927                            // running its MoE on the CPU.
10928                            tracing::warn!(
10929                                "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."
10930                            );
10931                            return None;
10932                        }
10933                        model.get_or_insert_with(|| mm.clone());
10934                        experts.push((gi, ui, di));
10935                    }
10936                    crate::gpu::GraphFfn::Moe {
10937                        router,
10938                        shared_gate: sgate,
10939                        experts,
10940                        n_exp: m.experts.len(),
10941                        // CMF_TOPK_PROBE: timing probe only — output is WRONG.
10942                        // Fewer experts shrink the MoE arithmetic while the
10943                        // dispatch count stays identical, which is the only
10944                        // clean way to tell a launch-bound decode from a
10945                        // compute-bound one.
10946                        top_k: std::env::var("CMF_TOPK_PROBE")
10947                            .ok()
10948                            .and_then(|v| v.parse::<usize>().ok())
10949                            .filter(|k| *k > 0 && *k <= m.top_k)
10950                            .unwrap_or(m.top_k),
10951                        inter,
10952                        norm_topk: m.norm_topk_prob,
10953                        q4tp: q4tp?,
10954                        gu_q2: gu_q2.unwrap_or(false),
10955                        sigmoid: m.router_sigmoid,
10956                        bias: m.expert_bias.as_deref(),
10957                        has_shared,
10958                        shared_gated,
10959                        route_scale: m.routed_scaling,
10960                    }
10961                }
10962            };
10963            let attn = match &lw.attn {
10964                AttnKind::Full {
10965                    wq,
10966                    wk,
10967                    wv,
10968                    wo,
10969                    q_norm,
10970                    k_norm,
10971                    output_gate,
10972                    softplus_gate,
10973                    bias,
10974                } => {
10975                    if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
10976                        return None;
10977                    }
10978                    let (m, _, _, _) = wq
10979                        .graph_weight()
10980                        .or_else(|| wq.graph_weight_descriptor())?;
10981                    model = Some(m.clone());
10982                    crate::gpu::GraphAttn::Full {
10983                        wq: gw(wq)?,
10984                        wk: gw(wk)?,
10985                        wv: gw(wv)?,
10986                        wo: gw(wo)?,
10987                        q_norm: q_norm.as_deref(),
10988                        k_norm: k_norm.as_deref(),
10989                        late_qk_norm: self.qk_norm_after_rope,
10990                        bias: bias
10991                            .as_ref()
10992                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
10993                        output_gate: *output_gate,
10994                        cpu_k: self.kv_cache.layers[li].k_heads(),
10995                        cpu_v: self.kv_cache.layers[li].v_heads(),
10996                        geom: self.graph_attn_geom(li),
10997                    }
10998                }
10999                AttnKind::LinearGdn(w) => {
11000                    let cfg = self.gdn_cfg?;
11001                    let (m, _, _, _) = w
11002                        .in_proj_qkv
11003                        .graph_weight()
11004                        .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
11005                    model = Some(m.clone());
11006                    crate::gpu::GraphAttn::Gdn {
11007                        qkv: gw(&w.in_proj_qkv)?,
11008                        z: gw(&w.in_proj_z)?,
11009                        a: gw(&w.in_proj_a)?,
11010                        b: gw(&w.in_proj_b)?,
11011                        out: gw(&w.out_proj)?,
11012                        conv1d: &w.conv1d,
11013                        a_log: &w.a_log,
11014                        dt_bias: &w.dt_bias,
11015                        norm: &w.norm,
11016                        nv: cfg.num_v_heads,
11017                        nk: cfg.num_k_heads,
11018                        dk: cfg.key_head_dim,
11019                        dv: cfg.value_head_dim,
11020                        kk: cfg.conv_kernel,
11021                        cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11022                    }
11023                }
11024                AttnKind::ShortConv(w) => {
11025                    let cfg = self.short_conv_cfg?;
11026                    let (m, _, _, _) = w
11027                        .in_proj
11028                        .graph_weight()
11029                        .or_else(|| w.in_proj.graph_weight_descriptor())?;
11030                    model = Some(m.clone());
11031                    crate::gpu::GraphAttn::ShortConv {
11032                        inp: gw(&w.in_proj)?,
11033                        out: gw(&w.out_proj)?,
11034                        taps: &w.conv,
11035                        kernel: cfg.kernel,
11036                        cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11037                    }
11038                }
11039                _ => return None,
11040            };
11041            layers.push(crate::gpu::GraphLayer {
11042                input_norm: &lw.input_norm,
11043                attn,
11044                post_norm: &lw.post_norm,
11045                ffn: gffn,
11046            });
11047        }
11048        let model = model?;
11049        // Fold final-norm + lm_head into the graph when this call wants logits
11050        // and the lm_head is a graphable (quantized) weight — the graph then
11051        // reads back logits (into logits_out) instead of the hidden, dropping
11052        // the separate CPU/GPU lm_head op + its sync. Never the f32 fallback:
11053        // an unquantized lm_head is vocab·hidden and must not be uploaded.
11054        let lm_gw = if upto_excl == self.num_layers
11055            && self.graph_want_logits
11056            && std::env::var("CMF_GPU_LMHEAD")
11057                .map(|v| v != "0")
11058                .unwrap_or(true)
11059        {
11060            self.weights
11061                .lm_head
11062                .graph_weight()
11063                .or_else(|| self.weights.lm_head.graph_weight_descriptor())
11064                .map(|(m, i, kind, rs)| {
11065                let name = &m.tensors[i].name;
11066                let prism = if crate::prism::is_inverse_embedding(m, name) {
11067                    crate::gpu::GraphPrismOp::InverseEmbedding
11068                } else if crate::prism::is_forward_weight(m, name) {
11069                    crate::gpu::GraphPrismOp::Forward
11070                } else {
11071                    crate::gpu::GraphPrismOp::None
11072                };
11073                (
11074                    crate::gpu::GraphW {
11075                        idx: i,
11076                        kind,
11077                        row_scale: rs,
11078                        data: &[],
11079                        prism,
11080                        affine: crate::prism::is_affine_target(m, name),
11081                    },
11082                    self.weights.lm_head.rows(),
11083                )
11084            })
11085        } else {
11086            None
11087        };
11088        let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
11089        // Multi-step re-embeds the winner on the device.
11090        let emb_gw = if steps > 1 {
11091            self.weights
11092                .embed_tokens
11093                .graph_weight()
11094                .or_else(|| self.weights.embed_tokens.graph_weight_descriptor())
11095                .map(|(m, i, kind, rs)| {
11096                    let name = &m.tensors[i].name;
11097                    let prism = if crate::prism::is_inverse_embedding(m, name) {
11098                        crate::gpu::GraphPrismOp::InverseEmbedding
11099                    } else if crate::prism::is_forward_weight(m, name) {
11100                        crate::gpu::GraphPrismOp::Forward
11101                    } else {
11102                        crate::gpu::GraphPrismOp::None
11103                    };
11104                    (
11105                        crate::gpu::GraphW {
11106                            idx: i,
11107                            kind,
11108                            row_scale: rs,
11109                            data: &[],
11110                            prism,
11111                            affine: crate::prism::is_affine_target(m, name),
11112                        },
11113                        self.weights.embed_tokens.rows(),
11114                        self.embed_multiplier,
11115                    )
11116                })
11117        } else {
11118            None
11119        };
11120
11121        // Loop boundaries: virtual layer indices after which final_norm is
11122        // applied (mid-stack only; the GLOBAL last layer's norm folds into
11123        // lm_head). Span-relative — the executor compares its enumerate
11124        // index. A span ending mid-stack keeps its boundary norm even when
11125        // it is the span's own last layer.
11126        let loop_norm_at: Vec<usize> = if self.loop_final_norm {
11127            (from..upto_excl.min(self.num_layers - 1))
11128                .filter(|&li| (li + 1) % self.physical_layers == 0)
11129                .map(|li| li - from)
11130                .collect()
11131        } else {
11132            Vec::new()
11133        };
11134        let mut h = hidden.to_vec();
11135        // The normal decode path only needs the fused lm-head logits.  A
11136        // CMF_LOGIT_DUMP diagnostic, however, promises a prompt-boundary
11137        // post-stack hidden alongside those logits; request the existing
11138        // second readback only for that explicit probe instead of dumping
11139        // the input copy left in `h` by a folded-head graph.
11140        let dump_hidden = std::env::var_os("CMF_LOGIT_DUMP").is_some();
11141        let outcome = crate::gpu::forward_token_graph(
11142            &model,
11143            self.graph_kv_id,
11144            &layers,
11145            &o1_views,
11146            self.o1_epoch,
11147            &self.inv_freq,
11148            &mut h,
11149            nh,
11150            nkv,
11151            hd,
11152            self.attn_scale,
11153            rd,
11154            self.hidden_size,
11155            self.intermediate_size,
11156            position,
11157            self.kv_cache.max_seq_len,
11158            gemma,
11159            self.rms_eps as f32,
11160            lm,
11161            &self.weights.final_norm,
11162            logits_out,
11163            &loop_norm_at,
11164            steps,
11165            emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
11166            ids_out,
11167            layers_run,
11168            from,
11169            dump_hidden,
11170        );
11171        match outcome {
11172            crate::gpu::TokenGraphOutcome::Completed => Some(Ok(h)),
11173            crate::gpu::TokenGraphOutcome::Failed => Some(Err(())),
11174            crate::gpu::TokenGraphOutcome::Declined => None,
11175        }
11176    }
11177
11178    /// Batched prefill: k contiguous prompt positions through the whole wgpu
11179    /// graph in ONE submit (projections/FFN as GEMMs). `hiddens` is [k·hidden]
11180    /// in/out (embeddings in, layer output out); KV mirror / GDN state advance.
11181    /// false ⇒ unsupported → caller keeps the per-position graph.
11182    /// The b-row Metal graph plan for the whole model: every layer as a
11183    /// GDN run or a full-attention item, all-or-nothing (a layer outside the
11184    /// graph's contract → None, the caller runs plain). Shared by the
11185    /// speculative verify and the batched prefill.
11186    #[cfg(target_os = "macos")]
11187    #[allow(clippy::type_complexity)]
11188    fn metal_rows_plan(
11189        &self,
11190    ) -> Option<(
11191        Vec<MetalRowsItem<'_>>,
11192        std::sync::Arc<cortiq_core::CmfModel>,
11193        Option<crate::gpu_metal::GdnGpuCfg>,
11194    )> {
11195        use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
11196        let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
11197        if !graph_force
11198            || !crate::gpu::enabled_here()
11199            || std::env::var("CMF_GPU_BLOCK")
11200                .map(|v| v == "0")
11201                .unwrap_or(false)
11202            || self.attn_softcap > 0.0
11203            || self.o1_active()
11204            || self.swa.is_some()
11205            || self.global_attn.is_some()
11206            || self.attention_heads_per_layer.is_some()
11207            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
11208            || self.graph_attn_decline_reason().is_some()
11209            || self.attn_v_norm
11210            || self.loop_final_norm
11211        {
11212            return None;
11213        }
11214        let attend_contract = self.head_dim % 4 == 0
11215            && self.head_dim <= 256
11216            && self.rotary_dim >= 2
11217            && self.rotary_dim <= self.head_dim
11218            && (self.rotary_dim / 2) % 32 == 0
11219            && self.num_kv_heads > 0
11220            && self.num_heads % self.num_kv_heads == 0;
11221        if !attend_contract {
11222            return None;
11223        }
11224        let mut plan: Vec<MetalRowsItem> = Vec::new();
11225        let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
11226        for li in 0..self.num_layers {
11227            let lw = &self.weights.layers[self.phys_layer(li)];
11228            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
11229                return None;
11230            }
11231            let ffn = match &lw.ffn {
11232                FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
11233                    let (Some(g), Some(u), Some(dn)) = (
11234                        d.gate_proj.metal_graph_parts(),
11235                        d.up_proj.metal_graph_parts(),
11236                        d.down_proj.metal_graph_parts(),
11237                    ) else {
11238                        return None;
11239                    };
11240                    MetalFfn::Dense {
11241                        gate: g,
11242                        up: u,
11243                        down: dn,
11244                    }
11245                }
11246                _ => return None,
11247            };
11248            match &lw.attn {
11249                AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
11250                    let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
11251                        w.in_proj_qkv.metal_graph_parts(),
11252                        w.in_proj_z.metal_graph_parts(),
11253                        w.in_proj_a.f32_parts(),
11254                        w.in_proj_b.f32_parts(),
11255                        w.out_proj.metal_graph_parts(),
11256                    ) else {
11257                        return None;
11258                    };
11259                    if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
11260                        model_ref.get_or_insert_with(|| model.clone());
11261                    }
11262                    let gl = GdnGpuLayer {
11263                        attn_norm: &lw.input_norm,
11264                        post_norm: &lw.post_norm,
11265                        qkv,
11266                        z,
11267                        a,
11268                        b: bb,
11269                        out,
11270                        ffn,
11271                        conv1d: &w.conv1d,
11272                        a_log: &w.a_log,
11273                        dt_bias: &w.dt_bias,
11274                        gnorm: &w.norm,
11275                    };
11276                    match plan.last_mut() {
11277                        Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
11278                        _ => plan.push(MetalRowsItem::Gdn {
11279                            run: vec![gl],
11280                            first: li,
11281                        }),
11282                    }
11283                }
11284                AttnKind::Full {
11285                    wq,
11286                    wk,
11287                    wv,
11288                    wo,
11289                    q_norm,
11290                    k_norm,
11291                    output_gate,
11292                    softplus_gate: None,
11293                    bias: None,
11294                } => {
11295                    let (Some(pq), Some(pk), Some(pv), Some(po)) =
11296                        (
11297                            wq.metal_graph_parts(),
11298                            wk.metal_graph_parts(),
11299                            wv.metal_graph_parts(),
11300                            wo.metal_graph_parts(),
11301                        )
11302                    else {
11303                        return None;
11304                    };
11305                    if let QTensor::Mapped { model, .. } = wq {
11306                        model_ref.get_or_insert_with(|| model.clone());
11307                    }
11308                    let cache = &self.kv_cache.layers[li];
11309                    if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
11310                        return None;
11311                    }
11312                    plan.push(MetalRowsItem::Attn {
11313                        l: AttnGpuLayer {
11314                            attn_norm: &lw.input_norm,
11315                            post_norm: &lw.post_norm,
11316                            wq: pq,
11317                            wk: pk,
11318                            wv: pv,
11319                            wo: po,
11320                            ffn,
11321                        },
11322                        li,
11323                        q_norm: q_norm.as_deref(),
11324                        k_norm: k_norm.as_deref(),
11325                        output_gate: *output_gate,
11326                    });
11327                }
11328                _ => return None,
11329            }
11330        }
11331        let model = model_ref?;
11332        let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
11333            nv: cfg.num_v_heads,
11334            nk: cfg.num_k_heads,
11335            dk: cfg.key_head_dim,
11336            dv: cfg.value_head_dim,
11337            kk: cfg.conv_kernel,
11338            hidden: self.hidden_size,
11339            inter: self.intermediate_size,
11340            c_dim: cfg.conv_dim(),
11341            eps: cfg.rms_eps as f32,
11342            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11343        });
11344        Some((plan, model, gcfg))
11345    }
11346
11347    /// `AttnDeviceParams` for a plan item over the CPU cache as it stands.
11348    #[cfg(target_os = "macos")]
11349    #[allow(clippy::too_many_arguments)]
11350    fn metal_attn_params<'a>(
11351        li: usize,
11352        cache: &'a crate::kv_cache::LayerKvCache,
11353        q_norm: Option<&'a [f32]>,
11354        k_norm: Option<&'a [f32]>,
11355        output_gate: bool,
11356        inv_freq: &'a [f32],
11357        geom: (usize, usize, usize, usize),
11358        pos0: usize,
11359        kv_id: u64,
11360        scale: f32,
11361        eps: f32,
11362        gemma: bool,
11363        late_qk_norm: bool,
11364    ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
11365        let (nh, nkv, hd, rd) = geom;
11366        let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11367        let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11368        let cpu_stored = cpu_k[0].len() / hd;
11369        (
11370            crate::gpu_metal::AttnDeviceParams {
11371                kv_id,
11372                layer: li,
11373                nh,
11374                nkv,
11375                hd,
11376                rd,
11377                position: pos0,
11378                scale,
11379                eps,
11380                gemma,
11381                late_qk_norm,
11382                output_gate,
11383                q_norm,
11384                k_norm,
11385                inv_freq,
11386                cpu_k,
11387                cpu_v,
11388                cpu_stored,
11389                o1: None,
11390            },
11391            cpu_stored,
11392        )
11393    }
11394
11395    /// Run the rows plan over `hiddens` (b rows at `pos0..`): validate,
11396    /// encode every item, optionally the head, sync. Returns the graph
11397    /// (for the commit / state finish) plus the GDN layer indices and the
11398    /// attention layers with the row count they were encoded against.
11399    #[cfg(target_os = "macos")]
11400    #[allow(clippy::type_complexity)]
11401    fn metal_rows_run(
11402        &mut self,
11403        hiddens: &mut [f32],
11404        pos0: usize,
11405        b: usize,
11406        prefill: bool,
11407        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11408        // Greedy verify: (row length scored, the b argmax ids out) — the
11409        // head's argmax runs on the device and the logits plane is NOT
11410        // read back (`spec.2` stays empty).
11411        mut argmax_out: Option<(usize, &mut Vec<u32>)>,
11412    ) -> MetalRowsRun {
11413        use crate::gpu_metal::{GraphDims, VerifyGraph};
11414        // The previous round's commit may still be replaying into the
11415        // trunk GDN owners on the second queue: this graph reads them
11416        // (zero-copy wraps) and may reallocate them below — collect the
11417        // replay first. Normally already complete (the draft chain ran
11418        // in between); a failed replay is terminal like a failed commit.
11419        if !crate::gpu_metal::wait_replay() {
11420            tracing::error!("Metal rows graph: the pending async replay failed");
11421            return MetalRowsRun::Failed;
11422        }
11423        spec_stamp("v.wait");
11424        // Seed the GDN recurrent records the rows graph reads — GDN layers
11425        // ONLY (the record the CPU path would allocate anyway).  Sizing
11426        // every layer here planted a zero GDN-sized record on a bounded
11427        // anchor of a natively bounded file this path then refused, and
11428        // that record counted as recurrent state on macOS.
11429        let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
11430        if want > 0 {
11431            let phys = self.physical_layers.max(1);
11432            for (li, l) in self.kv_cache.layers.iter_mut().enumerate() {
11433                let is_gdn = self
11434                    .weights
11435                    .layers
11436                    .get(li % phys)
11437                    .is_some_and(|lw| matches!(lw.attn, AttnKind::LinearGdn(_)));
11438                if is_gdn && l.linear_state.len() != want {
11439                    l.linear_state = vec![0f32; want];
11440                }
11441            }
11442        }
11443        let Some((plan, model, gcfg)) = self.metal_rows_plan() else {
11444            return MetalRowsRun::Declined;
11445        };
11446        spec_stamp("v.plan");
11447        let dims = GraphDims {
11448            hidden: self.hidden_size,
11449            eps: self.rms_eps as f32,
11450            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11451        };
11452        let Some(mut graph) = (if prefill {
11453            VerifyGraph::new_prefill(&model, dims, hiddens, b)
11454        } else {
11455            VerifyGraph::new(&model, dims, hiddens, b)
11456        }) else {
11457            return MetalRowsRun::Declined;
11458        };
11459        let geom = (
11460            self.num_heads,
11461            self.num_kv_heads,
11462            self.head_dim,
11463            self.rotary_dim,
11464        );
11465        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11466        let eps = self.rms_eps as f32;
11467        let kv_id = self.graph_kv_id;
11468        let inv_freq = self.inv_freq.clone();
11469        for item in &plan {
11470            let ok = match item {
11471                MetalRowsItem::Gdn { run, .. } => gcfg
11472                    .as_ref()
11473                    .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
11474                    .unwrap_or(false),
11475                MetalRowsItem::Attn {
11476                    l,
11477                    li,
11478                    q_norm,
11479                    k_norm,
11480                    output_gate,
11481                } => {
11482                    let (p, _) = Self::metal_attn_params(
11483                        *li,
11484                        &self.kv_cache.layers[*li],
11485                        *q_norm,
11486                        *k_norm,
11487                        *output_gate,
11488                        &inv_freq,
11489                        geom,
11490                        pos0,
11491                        kv_id,
11492                        self.attn_scale,
11493                        eps,
11494                        gemma,
11495                        self.qk_norm_after_rope,
11496                    );
11497                    graph.attn_ok(l, &p)
11498                }
11499            };
11500            if !ok {
11501                use std::sync::atomic::{AtomicBool, Ordering};
11502                static SAID: AtomicBool = AtomicBool::new(false);
11503                if !SAID.swap(true, Ordering::Relaxed) {
11504                    tracing::warn!("metal rows graph: a layer failed preflight — declining");
11505                }
11506                return MetalRowsRun::Declined;
11507            }
11508        }
11509        let lm = match &spec {
11510            Some((lm, _, _)) => {
11511                if !graph.lm_head_ok(*lm) {
11512                    return MetalRowsRun::Declined;
11513                }
11514                Some(*lm)
11515            }
11516            None => None,
11517        };
11518        let mut gdn_layers = Vec::new();
11519        let mut attn_layers = Vec::new();
11520        for item in &plan {
11521            match item {
11522                MetalRowsItem::Gdn { run, first } => {
11523                    let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
11524                        .iter()
11525                        .map(|l| l.linear_state.as_slice())
11526                        .collect();
11527                    if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
11528                        return MetalRowsRun::Declined;
11529                    }
11530                    gdn_layers.extend(*first..*first + run.len());
11531                }
11532                MetalRowsItem::Attn {
11533                    l,
11534                    li,
11535                    q_norm,
11536                    k_norm,
11537                    output_gate,
11538                } => {
11539                    let (p, cpu_stored) = Self::metal_attn_params(
11540                        *li,
11541                        &self.kv_cache.layers[*li],
11542                        *q_norm,
11543                        *k_norm,
11544                        *output_gate,
11545                        &inv_freq,
11546                        geom,
11547                        pos0,
11548                        kv_id,
11549                        self.attn_scale,
11550                        eps,
11551                        gemma,
11552                        self.qk_norm_after_rope,
11553                    );
11554                    if !graph.encode_attn_b(l, &p) {
11555                        return MetalRowsRun::Declined;
11556                    }
11557                    attn_layers.push((*li, cpu_stored));
11558                }
11559            }
11560        }
11561        if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
11562            if !graph.encode_lm_head_b(final_norm, lm) {
11563                return MetalRowsRun::Declined;
11564            }
11565            // The device argmax is an OPTIMISATION, never a reason to
11566            // decline the round: if it will not encode, drop it and read
11567            // the logits plane back the old way (the head is encoded
11568            // either way, so the rows are there).
11569            if let Some((n, _)) = argmax_out.as_ref() {
11570                if !graph.encode_argmax_b(*n) {
11571                    argmax_out = None;
11572                }
11573            }
11574        }
11575        spec_stamp("v.enc");
11576        if !graph.sync() {
11577            return MetalRowsRun::Failed;
11578        }
11579        spec_stamp("v.gpu");
11580        match (spec, argmax_out) {
11581            (Some(_), Some((_, ids))) => {
11582                ids.resize(b, 0);
11583                if !graph.read_argmax(ids) {
11584                    return MetalRowsRun::Failed;
11585                }
11586                spec_stamp("v.am");
11587            }
11588            (Some((lm, _, logits)), None) => {
11589                logits.resize(b * lm.1, 0.0);
11590                if !graph.read_logits(logits) {
11591                    return MetalRowsRun::Failed;
11592                }
11593                spec_stamp("v.lg");
11594            }
11595            (None, _) => {}
11596        }
11597        if !graph.read_hidden(hiddens) {
11598            return MetalRowsRun::Failed;
11599        }
11600        spec_stamp("v.hid");
11601        MetalRowsRun::Completed(MetalVerifyPending {
11602            graph,
11603            gdn_layers,
11604            attn_layers,
11605        })
11606    }
11607
11608    /// Native-Metal twin of `try_batch_graph_wgpu`: the b rows through the
11609    /// whole model on the `VerifyGraph` (one submit), the head folded in
11610    /// when `spec` asks; `hiddens` come back as the last layer's output
11611    /// rows, `spec.2` as `[b][lm_rows]` logits. The graph is parked in
11612    /// `metal_verify` for `metal_verify_commit`.
11613    #[cfg(target_os = "macos")]
11614    fn try_batch_graph_metal(
11615        &mut self,
11616        hiddens: &mut [f32],
11617        positions: &[usize],
11618        b: usize,
11619        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11620        argmax_out: Option<(usize, &mut Vec<u32>)>,
11621    ) -> crate::gpu::BatchGraphOutcome {
11622        let _t0 = std::time::Instant::now();
11623        if positions.len() != b
11624            || positions.windows(2).any(|w| w[1] != w[0] + 1)
11625            || hiddens.len() != b * self.hidden_size
11626        {
11627            return crate::gpu::BatchGraphOutcome::Declined;
11628        }
11629        let pending = match self.metal_rows_run(hiddens, positions[0], b, false, spec, argmax_out) {
11630            MetalRowsRun::Declined => return crate::gpu::BatchGraphOutcome::Declined,
11631            MetalRowsRun::Failed => return crate::gpu::BatchGraphOutcome::Failed,
11632            MetalRowsRun::Completed(pending) => pending,
11633        };
11634        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
11635            eprintln!(
11636                "metal-verify: {:.1} ms | b={b}",
11637                _t0.elapsed().as_secs_f64() * 1e3
11638            );
11639        }
11640        self.metal_verify = Some(pending);
11641        crate::gpu::BatchGraphOutcome::Completed
11642    }
11643
11644    /// Batched prefill on the Metal rows graph: `ids` (≤ 512) at
11645    /// `start_pos..`, states written in place, K/V rows appended to the
11646    /// CPU caches; optional final norm/head logits are returned in `spec`.
11647    /// Declined means no command buffer was admitted; Failed is terminal.
11648    #[cfg(target_os = "macos")]
11649    fn prefill_rows_metal(
11650        &mut self,
11651        ids: &[u32],
11652        start_pos: usize,
11653        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11654    ) -> MetalPrefillOutcome {
11655        let b = ids.len();
11656        if b == 0 || b > 512 {
11657            return MetalPrefillOutcome::Declined;
11658        }
11659        METAL_PREFILL_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11660        let with_head = spec.is_some();
11661        let hs = self.hidden_size;
11662        let mut hiddens = vec![0f32; b * hs];
11663        for (j, &id) in ids.iter().enumerate() {
11664            let e = self.embed_single(id);
11665            hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
11666        }
11667        let mut pending = match self.metal_rows_run(&mut hiddens, start_pos, b, true, spec, None) {
11668            MetalRowsRun::Declined => return MetalPrefillOutcome::Declined,
11669            MetalRowsRun::Failed => {
11670                METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11671                return MetalPrefillOutcome::Failed;
11672            }
11673            MetalRowsRun::Completed(pending) => pending,
11674        };
11675        // states are final: copy them to the owners
11676        let idxs = pending.gdn_layers.clone();
11677        let mut outs: Vec<&mut [f32]> = self
11678            .kv_cache
11679            .layers
11680            .iter_mut()
11681            .enumerate()
11682            .filter(|(i, _)| idxs.binary_search(i).is_ok())
11683            .map(|(_, l)| l.linear_state.as_mut_slice())
11684            .collect();
11685        if !pending.graph.finish_states(&mut outs) {
11686            METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11687            return MetalPrefillOutcome::Failed;
11688        }
11689        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
11690        // Read every layer before mutating any CPU cache.  A missing mirror
11691        // row is a terminal graph failure, not a reason to append a partial
11692        // prefix and replay the remainder serially.
11693        let mut rows = Vec::with_capacity(pending.attn_layers.len());
11694        for (li, cpu_stored) in &pending.attn_layers {
11695            let mut kbuf = vec![0f32; b * nkv * hd];
11696            let mut vbuf = vec![0f32; b * nkv * hd];
11697            if !crate::gpu_metal::kv_mirror_read_rows(
11698                self.graph_kv_id,
11699                *li,
11700                nkv,
11701                hd,
11702                *cpu_stored,
11703                b,
11704                &mut kbuf,
11705                &mut vbuf,
11706            ) {
11707                METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
11708                return MetalPrefillOutcome::Failed;
11709            }
11710            rows.push((*li, *cpu_stored, kbuf, vbuf));
11711        }
11712        for (li, cpu_stored, kbuf, vbuf) in rows {
11713            let cache = &mut self.kv_cache.layers[li];
11714            for r in 0..b {
11715                cache.append(
11716                    &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
11717                    &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
11718                    &[],
11719                );
11720            }
11721            crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + b);
11722        }
11723        METAL_PREFILL_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11724        if with_head {
11725            METAL_PREFILL_HEAD_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
11726        }
11727        MetalPrefillOutcome::Completed(hiddens)
11728    }
11729
11730    #[cfg(target_os = "macos")]
11731    fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> MetalPrefillOutcome {
11732        self.prefill_rows_metal(ids, start_pos, None)
11733    }
11734
11735    /// Exact teacher-forced NLL through the ordinary Metal rows graph.  This
11736    /// is intentionally separate from the serial TokenGraph scorer: every
11737    /// chunk owns a real b-row graph/head completion and the recurrent/KV
11738    /// handoff is committed before the next chunk begins.
11739    #[cfg(target_os = "macos")]
11740    fn nll_batch_metal(&mut self, ids: &[u32], start: usize) -> MetalBatchNllOutcome {
11741        if ids.len() < 2 || self.o1_active() || self.head_clusters.is_some() {
11742            return MetalBatchNllOutcome::Declined;
11743        }
11744        let Some(lm) = self.weights.lm_head.metal_graph_parts() else {
11745            return MetalBatchNllOutcome::Declined;
11746        };
11747        let chunk = std::env::var("CMF_METAL_PREFILL_CHUNK")
11748            .ok()
11749            .and_then(|v| v.parse::<usize>().ok())
11750            .filter(|&v| (1..=512).contains(&v))
11751            .unwrap_or(32);
11752        let final_norm = self.weights.final_norm.clone();
11753        let mut nll = 0.0f64;
11754        let mut count = 0usize;
11755        let mut pos = 0usize;
11756        let mut completed = 0usize;
11757        while pos < ids.len() {
11758            let end = (pos + chunk).min(ids.len());
11759            let mut logits = Vec::new();
11760            let outcome = self.prefill_rows_metal(
11761                &ids[pos..end],
11762                pos,
11763                Some((lm, &final_norm, &mut logits)),
11764            );
11765            match outcome {
11766                MetalPrefillOutcome::Declined => {
11767                    return if completed == 0 {
11768                        MetalBatchNllOutcome::Declined
11769                    } else {
11770                        MetalBatchNllOutcome::Failed(format!(
11771                            "ordinary Metal NLL batch declined after {completed} chunks"
11772                        ))
11773                    };
11774                }
11775                MetalPrefillOutcome::Failed => {
11776                    return MetalBatchNllOutcome::Failed(
11777                        "ordinary Metal NLL batch failed after admission".to_string(),
11778                    );
11779                }
11780                MetalPrefillOutcome::Completed(_) => {}
11781            }
11782            completed += 1;
11783            let vocab = self.vocab_size.min(lm.1);
11784            if logits.len() != (end - pos) * lm.1 || vocab == 0 {
11785                return MetalBatchNllOutcome::Failed(
11786                    "ordinary Metal NLL head returned an invalid shape".to_string(),
11787                );
11788            }
11789            for row in 0..(end - pos) {
11790                let absolute = pos + row;
11791                if absolute < start || absolute + 1 >= ids.len() {
11792                    continue;
11793                }
11794                let lg = &mut logits[row * lm.1..row * lm.1 + vocab];
11795                if let Some(mu) = self.logit_multiplier {
11796                    for v in lg.iter_mut() {
11797                        *v *= mu;
11798                    }
11799                }
11800                if let Some(c) = self.final_softcap {
11801                    for v in lg.iter_mut() {
11802                        *v = c * (*v / c).tanh();
11803                    }
11804                }
11805                let target = ids[absolute + 1] as usize;
11806                if target >= vocab {
11807                    return MetalBatchNllOutcome::Failed(format!(
11808                        "target token {target} exceeds Metal head rows {vocab}"
11809                    ));
11810                }
11811                let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
11812                let lse: f64 = lg
11813                    .iter()
11814                    .map(|&v| ((v - max) as f64).exp())
11815                    .sum::<f64>()
11816                    .ln()
11817                    + max as f64;
11818                nll += lse - lg[target] as f64;
11819                count += 1;
11820            }
11821            pos = end;
11822        }
11823        MetalBatchNllOutcome::Completed(nll, count)
11824    }
11825
11826    /// Commit a Metal verify round: replay the GDN recurrences over the
11827    /// `a + 1` accepted positions into the CPU states, append the accepted
11828    /// K/V rows from the mirrors to the CPU caches, re-point the mirrors.
11829    #[cfg(target_os = "macos")]
11830    fn metal_verify_commit(&mut self, a: usize) -> bool {
11831        let Some(mut pending) = self.metal_verify.take() else {
11832            return false;
11833        };
11834        let n = a + 1;
11835        // encode order == ascending layer order (the plan walks 0..layers)
11836        let idxs = pending.gdn_layers.clone();
11837        let mut outs: Vec<&mut [f32]> = self
11838            .kv_cache
11839            .layers
11840            .iter_mut()
11841            .enumerate()
11842            .filter(|(i, _)| idxs.binary_search(i).is_ok())
11843            .map(|(_, l)| l.linear_state.as_mut_slice())
11844            .collect();
11845        if !pending.graph.commit(n, &mut outs) {
11846            return false;
11847        }
11848        spec_stamp("c.replay");
11849        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
11850        // Read every layer before mutating any CPU cache.  Missing rows are
11851        // terminal after the replay has executed; never append a partial KV
11852        // prefix and continue on a serial path.
11853        let mut rows = Vec::with_capacity(pending.attn_layers.len());
11854        for (li, cpu_stored) in &pending.attn_layers {
11855            let mut kbuf = vec![0f32; n * nkv * hd];
11856            let mut vbuf = vec![0f32; n * nkv * hd];
11857            if !crate::gpu_metal::kv_mirror_read_rows(
11858                self.graph_kv_id,
11859                *li,
11860                nkv,
11861                hd,
11862                *cpu_stored,
11863                n,
11864                &mut kbuf,
11865                &mut vbuf,
11866            ) {
11867                return false;
11868            }
11869            rows.push((*li, *cpu_stored, kbuf, vbuf));
11870        }
11871        for (li, cpu_stored, kbuf, vbuf) in rows {
11872            let cache = &mut self.kv_cache.layers[li];
11873            for r in 0..n {
11874                cache.append(
11875                    &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
11876                    &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
11877                    &[],
11878                );
11879            }
11880            crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + n);
11881        }
11882        spec_stamp("c.kv");
11883        true
11884    }
11885
11886    /// The round's warm-ups as ONE b-row graph run over the MTP block on
11887    /// Metal: `pairs` = (trunk hidden, next token) at consecutive positions
11888    /// from `first_pos`; the block's input projection is folded in. This
11889    /// half encodes and SUBMITS (no wait); `mtp_warm_batch_finish` waits
11890    /// and pulls the appended K/V rows into the CPU MTP cache. None = the
11891    /// graph declined (nothing submitted, nothing appended).
11892    #[cfg(target_os = "macos")]
11893    fn mtp_warm_batch_submit(
11894        &mut self,
11895        m: &mut MtpModule,
11896        pairs: &[(&[f32], u32)],
11897        first_pos: usize,
11898    ) -> Option<MetalWarmPending> {
11899        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
11900        let b = pairs.len();
11901        if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
11902            return None;
11903        }
11904        let AttnKind::Full {
11905            wq,
11906            wk,
11907            wv,
11908            wo,
11909            q_norm,
11910            k_norm,
11911            output_gate,
11912            softplus_gate: None,
11913            bias: None,
11914        } = &m.layer.attn
11915        else {
11916            return None;
11917        };
11918        let FfnKind::Dense(d) = &m.layer.ffn else {
11919            return None;
11920        };
11921        if !d.segs.is_empty() {
11922            return None;
11923        }
11924        let (Some(pq), Some(pk), Some(pv), Some(po)) =
11925            (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
11926        else {
11927            return None;
11928        };
11929        let (Some(g), Some(u), Some(dn)) = (
11930            d.gate_proj.q1_parts(),
11931            d.up_proj.q1_parts(),
11932            d.down_proj.q1_parts(),
11933        ) else {
11934            return None;
11935        };
11936        let Some(eh) = m.eh_proj.q1_parts() else {
11937            return None;
11938        };
11939        let QTensor::Mapped { model, .. } = wq else {
11940            return None;
11941        };
11942        let model = model.clone();
11943        let hs = self.hidden_size;
11944        // [enorm(embed(tok)); hnorm(hidden)] rows
11945        let mut cat = vec![0f32; b * 2 * hs];
11946        for (j, (h, tok)) in pairs.iter().enumerate() {
11947            let e = self.embed_single(*tok);
11948            let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
11949            inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
11950            inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
11951        }
11952        let dims = GraphDims {
11953            hidden: hs,
11954            eps: self.rms_eps as f32,
11955            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11956        };
11957        spec_stamp("w.cat");
11958        let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
11959            return None;
11960        };
11961        spec_stamp("w.new");
11962        let l = AttnGpuLayer {
11963            attn_norm: &m.layer.input_norm,
11964            post_norm: &m.layer.post_norm,
11965            wq: pq,
11966            wk: pk,
11967            wv: pv,
11968            wo: po,
11969            ffn: MetalFfn::Dense {
11970                gate: g,
11971                up: u,
11972                down: dn,
11973            },
11974        };
11975        let (nh, nkv, hd, rd) = (
11976            self.num_heads,
11977            self.num_kv_heads,
11978            self.head_dim,
11979            self.rotary_dim,
11980        );
11981        let inv_freq = self.inv_freq.clone();
11982        let cpu_stored;
11983        {
11984            let cache = &m.kv;
11985            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11986            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11987            cpu_stored = cpu_k[0].len() / hd;
11988            // The cache may LAG the position (rows nobody warmed): the
11989            // pairs land at cpu_stored.. with their true RoPE positions
11990            // first_pos.., exactly what the one-by-one warm does. A cache
11991            // AHEAD of the position is a real inconsistency.
11992            if cpu_stored > first_pos {
11993                spec_stamp("w.decl");
11994                return None;
11995            }
11996            let p = AttnDeviceParams {
11997                kv_id: self.mtp_kv_id(),
11998                layer: Self::MTP_LAYER_BASE,
11999                nh,
12000                nkv,
12001                hd,
12002                rd,
12003                position: first_pos,
12004                scale: self.attn_scale,
12005                eps: self.rms_eps as f32,
12006                gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12007                late_qk_norm: self.qk_norm_after_rope,
12008                output_gate: *output_gate,
12009                q_norm: q_norm.as_deref(),
12010                k_norm: k_norm.as_deref(),
12011                inv_freq: &inv_freq,
12012                cpu_k,
12013                cpu_v,
12014                cpu_stored,
12015                o1: None,
12016            };
12017            if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
12018                return None;
12019            }
12020        }
12021        spec_stamp("w.enc");
12022        if !graph.submit() {
12023            return None;
12024        }
12025        spec_stamp("w.sub");
12026        Some(MetalWarmPending {
12027            graph,
12028            cpu_stored,
12029            b,
12030        })
12031    }
12032
12033    /// Submit and finish in one call (the prefill's MTP warm-up, where
12034    /// nothing runs in between).
12035    #[cfg(target_os = "macos")]
12036    fn mtp_warm_batch_metal(
12037        &mut self,
12038        m: &mut MtpModule,
12039        pairs: &[(&[f32], u32)],
12040        first_pos: usize,
12041    ) -> bool {
12042        match self.mtp_warm_batch_submit(m, pairs, first_pos) {
12043            Some(p) => self.mtp_warm_batch_finish(m, p),
12044            None => false,
12045        }
12046    }
12047
12048    /// Second half of the batched warm-up: wait for the submitted graph,
12049    /// pull its b appended K/V rows into the CPU MTP cache, re-point the
12050    /// mirror. False = the command buffer failed or the rows are missing
12051    /// (nothing appended; the caller falls back to the one-by-one warm).
12052    #[cfg(target_os = "macos")]
12053    fn mtp_warm_batch_finish(&mut self, m: &mut MtpModule, pending: MetalWarmPending) -> bool {
12054        let MetalWarmPending {
12055            mut graph,
12056            cpu_stored,
12057            b,
12058        } = pending;
12059        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12060        if !graph.sync() {
12061            return false;
12062        }
12063        spec_stamp("w.gpu");
12064        let mut kbuf = vec![0f32; b * nkv * hd];
12065        let mut vbuf = vec![0f32; b * nkv * hd];
12066        if !crate::gpu_metal::kv_mirror_read_rows(
12067            self.mtp_kv_id(),
12068            Self::MTP_LAYER_BASE,
12069            nkv,
12070            hd,
12071            cpu_stored,
12072            b,
12073            &mut kbuf,
12074            &mut vbuf,
12075        ) {
12076            return false;
12077        }
12078        for r in 0..b {
12079            m.kv.append(
12080                &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12081                &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12082                &[],
12083            );
12084        }
12085        crate::gpu_metal::kv_mirror_set_stored(
12086            self.mtp_kv_id(),
12087            Self::MTP_LAYER_BASE,
12088            cpu_stored + b,
12089        );
12090        spec_stamp("w.kv");
12091        true
12092    }
12093
12094    /// A committed token id from the high table (Cyrillic, CJK and the
12095    /// like sit above 131072 in Qwen's vocabulary; Latin subwords past
12096    /// the 65536 cut are rare enough to lose as rejected drafts) switches
12097    /// the draft to the full head for the next 16 tokens; other ids count
12098    /// down. On an M4 the full 660 MB head costs 5.5 ms a draft step
12099    /// against 1.4 for the shortlist, so the streak is kept short.
12100    pub(crate) fn note_draft_id(&mut self, id: u32) {
12101        let cut = Self::draft_vocab_rows(usize::MAX).max(131_072);
12102        if (id as usize) >= cut {
12103            self.draft_full_streak = 16;
12104        } else {
12105            self.draft_full_streak = self.draft_full_streak.saturating_sub(1);
12106        }
12107    }
12108
12109    /// The draft head's rows for the next step: the shortlist, or the full
12110    /// head while `draft_full_streak` runs.
12111    fn draft_head_rows(&self, head_rows: usize) -> usize {
12112        if self.draft_full_streak > 0 {
12113            head_rows
12114        } else {
12115            Self::draft_vocab_rows(head_rows)
12116        }
12117    }
12118
12119    /// Draft-head shortlist size: `CMF_DRAFT_VOCAB` rows (default 65536,
12120    /// capped at the head; 0 = full head).
12121    fn draft_vocab_rows(head_rows: usize) -> usize {
12122        static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
12123        let n = *N.get_or_init(|| {
12124            std::env::var("CMF_DRAFT_VOCAB")
12125                .ok()
12126                .and_then(|v| v.parse().ok())
12127                .unwrap_or(65536)
12128        });
12129        if n == 0 { head_rows } else { n.min(head_rows) }
12130    }
12131
12132    /// One MTP block step on the native Metal token graph: block input on
12133    /// the host, the attention layer + FFN device-resident over the MTP
12134    /// mirror, the head folded in when `want_logits`. The appended K/V row
12135    /// is pulled into the CPU MTP cache (owner of record) after the sync.
12136    #[cfg(target_os = "macos")]
12137    fn mtp_step_metal(
12138        &mut self,
12139        m: &mut MtpModule,
12140        hidden: &[f32],
12141        next_token: u32,
12142        position: usize,
12143        want_logits: bool,
12144    ) -> Option<(Vec<f32>, Vec<f32>)> {
12145        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12146        if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12147            || !crate::gpu::q1_force()
12148            || !crate::gpu::enabled_here()
12149            || self.attn_softcap > 0.0
12150            || self.attention_heads_per_layer.is_some()
12151            || m.kv.mode != crate::kv_cache::KvMode::F32
12152            || m.kv.o1.is_some()
12153        {
12154            return None;
12155        }
12156        let AttnKind::Full {
12157            wq,
12158            wk,
12159            wv,
12160            wo,
12161            q_norm,
12162            k_norm,
12163            output_gate,
12164            softplus_gate: None,
12165            bias: None,
12166        } = &m.layer.attn
12167        else {
12168            return None;
12169        };
12170        let FfnKind::Dense(d) = &m.layer.ffn else {
12171            return None;
12172        };
12173        if d.act != Act::Silu || !d.segs.is_empty() {
12174            return None;
12175        }
12176        let (pq, pk, pv, po) = (
12177            wq.q1_parts()?,
12178            wk.q1_parts()?,
12179            wv.q1_parts()?,
12180            wo.q1_parts()?,
12181        );
12182        let (g, u, dn) = (
12183            d.gate_proj.q1_parts()?,
12184            d.up_proj.q1_parts()?,
12185            d.down_proj.q1_parts()?,
12186        );
12187        let QTensor::Mapped { model, .. } = wq else {
12188            return None;
12189        };
12190        let model = model.clone();
12191        let lm = if want_logits {
12192            Some(self.weights.lm_head.q1_parts()?)
12193        } else {
12194            None
12195        };
12196        let dims = GraphDims {
12197            hidden: self.hidden_size,
12198            eps: self.rms_eps as f32,
12199            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12200        };
12201        // The block input `eh_proj · [enorm(e); hnorm(h)]` rides in the
12202        // graph (one submit a step); the host per-op matvec if it cannot.
12203        let hs = self.hidden_size;
12204        let mut x = vec![0f32; hs];
12205        let mut graph = TokenGraph::new(&model, dims, &x)?;
12206        let mut folded = false;
12207        if let Some(eh) = m.eh_proj.q1_parts() {
12208            let e = self.embed_single(next_token);
12209            let mut cat = vec![0.0f32; 2 * hs];
12210            let (cat_e, cat_h) = cat.split_at_mut(hs);
12211            inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
12212            inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
12213            folded = graph.encode_input_proj(eh, &cat);
12214        }
12215        if !folded {
12216            x = self.mtp_block_input(m, hidden, next_token);
12217            graph = TokenGraph::new(&model, dims, &x)?;
12218        }
12219        spec_stamp("d.in");
12220        let l = AttnGpuLayer {
12221            attn_norm: &m.layer.input_norm,
12222            post_norm: &m.layer.post_norm,
12223            wq: pq,
12224            wk: pk,
12225            wv: pv,
12226            wo: po,
12227            ffn: MetalFfn::Dense {
12228                gate: g,
12229                up: u,
12230                down: dn,
12231            },
12232        };
12233        let (nh, nkv, hd, rd) = (
12234            self.num_heads,
12235            self.num_kv_heads,
12236            self.head_dim,
12237            self.rotary_dim,
12238        );
12239        let inv_freq = self.inv_freq.clone();
12240        {
12241            let cache = &m.kv;
12242            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12243            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12244            let cpu_stored = cpu_k[0].len() / hd;
12245            let p = AttnDeviceParams {
12246                kv_id: self.mtp_kv_id(),
12247                layer: Self::MTP_LAYER_BASE,
12248                nh,
12249                nkv,
12250                hd,
12251                rd,
12252                position,
12253                scale: self.attn_scale,
12254                eps: self.rms_eps as f32,
12255                gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12256                late_qk_norm: self.qk_norm_after_rope,
12257                output_gate: *output_gate,
12258                q_norm: q_norm.as_deref(),
12259                k_norm: k_norm.as_deref(),
12260                inv_freq: &inv_freq,
12261                cpu_k,
12262                cpu_v,
12263                cpu_stored,
12264                o1: None,
12265            };
12266            if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12267                return None;
12268            }
12269        }
12270        // The draft's head over a vocabulary SHORTLIST (the first
12271        // CMF_DRAFT_VOCAB rows — BPE ids run roughly by merge rank, so the
12272        // low ids carry the mass): the verify keeps the full head, so a true
12273        // token past the cut is only a rejected draft, never a wrong token.
12274        // 662 MB a step on Qwen3.8 becomes 170 MB at 65536.
12275        let draft_rows = if let Some(lm) = lm {
12276            self.draft_head_rows(lm.1)
12277        } else {
12278            0
12279        };
12280        if let Some(lm) = lm {
12281            if !graph.lm_head_ok(lm) {
12282                return None;
12283            }
12284            if draft_rows < lm.1 {
12285                if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12286                    return None;
12287                }
12288            } else {
12289                graph.encode_lm_head(&m.final_norm, lm);
12290            }
12291        }
12292        spec_stamp("d.enc");
12293        if graph.sync_checked().is_err() {
12294            return None;
12295        }
12296        spec_stamp("d.gpu");
12297        let mut logits = Vec::new();
12298        if let Some(lm) = lm {
12299            let n_read = draft_rows.min(lm.1).min(self.vocab_size);
12300            logits = attention::take_buf(n_read);
12301            graph.read_logits(&mut logits);
12302            // ids past the shortlist: never drafted (−∞ in every chain)
12303            logits.resize(self.vocab_size, f32::NEG_INFINITY);
12304        }
12305        graph.finish(&mut x);
12306        let mut krow = attention::take_buf(nkv * hd);
12307        let mut vrow = attention::take_buf(nkv * hd);
12308        if crate::gpu_metal::kv_mirror_read_last(
12309            self.mtp_kv_id(),
12310            Self::MTP_LAYER_BASE,
12311            nkv,
12312            hd,
12313            &mut krow,
12314            &mut vrow,
12315        ) {
12316            m.kv.append(&krow, &vrow, &[]);
12317        }
12318        attention::recycle_buf(&mut krow);
12319        attention::recycle_buf(&mut vrow);
12320        spec_stamp("d.rd");
12321        Some((logits, x))
12322    }
12323
12324    /// `CMF_MTP_CHAIN=0` keeps the per-step draft (one submit and one
12325    /// host round trip per MTP step); the default drafts the whole chain
12326    /// in one command buffer when the round is plain greedy.
12327    ///
12328    /// Measured on an M4 (24 GB), Qwen3.8-27B q4tp, P3 at 160 tokens,
12329    /// k=7, six runs per arm alternating inside one lock window — the
12330    /// round's draft phase (median over the 34 rounds of a run) is
12331    /// 34.5 ms per round old against 30.1 new, i.e. 4.93 → 4.31 ms per
12332    /// draft step. That is the whole prize: the 7 submits cost ~0.6 ms
12333    /// each in host and submit latency and nothing else changes —
12334    /// acceptance (3.41 of 7) and tokens per round (4.41) are identical,
12335    /// and the round is 289 → 285 ms, decode 13.8 → 14.0 tok/s.
12336    fn mtp_chain_on() -> bool {
12337        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12338        *ON.get_or_init(|| std::env::var("CMF_MTP_CHAIN").as_deref() != Ok("0"))
12339    }
12340
12341    /// The round's k greedy drafts as ONE command buffer on Metal: the MTP
12342    /// block k times back to back, each step's token embedding gathered
12343    /// on the device from the argmax the step before it wrote, the head
12344    /// over the round's shortlist (or the full head during a full-head
12345    /// streak — decided once, before the chain, exactly as the per-step
12346    /// path decides it per step, since `draft_full_streak` only moves on
12347    /// a commit). One wait, then the k ids and the k appended K/V rows
12348    /// come back; the CPU MTP cache ends where k `mtp_step_metal` calls
12349    /// would have left it. `Err(false)` = declined before anything was
12350    /// committed (the per-step path takes the round); `Err(true)` = the
12351    /// command buffer failed after commit.
12352    #[cfg(target_os = "macos")]
12353    fn mtp_draft_chain_metal(
12354        &mut self,
12355        m: &mut MtpModule,
12356        hidden: &[f32],
12357        t_next: u32,
12358        position: usize,
12359        k: usize,
12360    ) -> Result<Vec<u32>, bool> {
12361        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12362        if k == 0
12363            || k > 64
12364            || !Self::mtp_chain_on()
12365            || std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12366            || !crate::gpu::q1_force()
12367            || !crate::gpu::enabled_here()
12368            || self.attn_softcap > 0.0
12369            || self.attention_heads_per_layer.is_some()
12370            || m.kv.mode != crate::kv_cache::KvMode::F32
12371            || m.kv.o1.is_some()
12372            // the chain gathers embeddings itself: only the plain table
12373            || self.dsv4.is_some()
12374            || self.dsv41.is_some()
12375            || self.qwen4_exp.is_some()
12376            || self.g3n.is_some()
12377        {
12378            return Err(false);
12379        }
12380        let AttnKind::Full {
12381            wq,
12382            wk,
12383            wv,
12384            wo,
12385            q_norm,
12386            k_norm,
12387            output_gate,
12388            softplus_gate: None,
12389            bias: None,
12390        } = &m.layer.attn
12391        else {
12392            return Err(false);
12393        };
12394        let FfnKind::Dense(d) = &m.layer.ffn else {
12395            return Err(false);
12396        };
12397        if d.act != Act::Silu || !d.segs.is_empty() {
12398            return Err(false);
12399        }
12400        let (Some(pq), Some(pk), Some(pv), Some(po)) =
12401            (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12402        else {
12403            return Err(false);
12404        };
12405        let (Some(g), Some(u), Some(dn)) = (
12406            d.gate_proj.q1_parts(),
12407            d.up_proj.q1_parts(),
12408            d.down_proj.q1_parts(),
12409        ) else {
12410            return Err(false);
12411        };
12412        let (Some(eh), Some(lm)) = (m.eh_proj.q1_parts(), self.weights.lm_head.q1_parts()) else {
12413            return Err(false);
12414        };
12415        let QTensor::Mapped { model, .. } = wq else {
12416            return Err(false);
12417        };
12418        let model = model.clone();
12419        // the embedding table: a q4tp tensor of the SAME blob, no Prism
12420        // inverse-embedding post-pass
12421        let QTensor::Mapped {
12422            model: em,
12423            idx: eidx,
12424            dtype: cortiq_core::TensorDtype::Q4TiledP,
12425            ..
12426        } = &self.weights.embed_tokens
12427        else {
12428            return Err(false);
12429        };
12430        if !std::sync::Arc::ptr_eq(em, &model)
12431            || crate::prism::is_inverse_embedding(&model, &model.tensors[*eidx].name)
12432        {
12433            return Err(false);
12434        }
12435        let embed = (
12436            *eidx,
12437            self.weights.embed_tokens.rows(),
12438            self.weights.embed_tokens.cols(),
12439        );
12440        if embed.2 != self.hidden_size || hidden.len() != self.hidden_size {
12441            return Err(false);
12442        }
12443        let dims = GraphDims {
12444            hidden: self.hidden_size,
12445            eps: self.rms_eps as f32,
12446            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12447        };
12448        let Some(mut graph) = TokenGraph::new(&model, dims, hidden) else {
12449            return Err(false);
12450        };
12451        if !graph.chain_embed_ok(embed) || !graph.lm_head_ok(lm) {
12452            return Err(false);
12453        }
12454        let l = AttnGpuLayer {
12455            attn_norm: &m.layer.input_norm,
12456            post_norm: &m.layer.post_norm,
12457            wq: pq,
12458            wk: pk,
12459            wv: pv,
12460            wo: po,
12461            ffn: MetalFfn::Dense {
12462                gate: g,
12463                up: u,
12464                down: dn,
12465            },
12466        };
12467        let (nh, nkv, hd, rd) = (
12468            self.num_heads,
12469            self.num_kv_heads,
12470            self.head_dim,
12471            self.rotary_dim,
12472        );
12473        let inv_freq = self.inv_freq.clone();
12474        let draft_rows = self.draft_head_rows(lm.1);
12475        let n_arg = draft_rows.min(lm.1).min(self.vocab_size);
12476        if n_arg == 0 {
12477            return Err(false);
12478        }
12479        // `CMF_MTP_CHAIN_SPLIT=1` commits each step as it is encoded, so
12480        // the GPU starts on step 0 while the host is still encoding step
12481        // 1 — a probe for whether the host encode is on the critical
12482        // path. It is not: three runs each, draft 30.0 ms per round split
12483        // against 30.1 whole, and the whole chain's host encode measures
12484        // 0.3 ms against a 29.7 ms wait. Kept as a probe, off by default.
12485        let split = std::env::var("CMF_MTP_CHAIN_SPLIT").as_deref() == Ok("1");
12486        let t_chain = std::time::Instant::now();
12487        graph.chain_ids_init(t_next, k);
12488        let cpu_stored;
12489        {
12490            let cache = &m.kv;
12491            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12492            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12493            cpu_stored = cpu_k[0].len() / hd;
12494            for j in 0..k {
12495                if !graph.encode_chain_input(
12496                    embed,
12497                    j as u32,
12498                    &m.enorm,
12499                    &m.hnorm,
12500                    self.embed_multiplier,
12501                    eh,
12502                ) {
12503                    return Err(false);
12504                }
12505                // step j's mirror row: the mirror is re-pointed at the CPU
12506                // rows before step 0 and advances by one per step; its
12507                // resync (never taken past step 0) reads the CPU rows
12508                let p = AttnDeviceParams {
12509                    kv_id: self.mtp_kv_id(),
12510                    layer: Self::MTP_LAYER_BASE,
12511                    nh,
12512                    nkv,
12513                    hd,
12514                    rd,
12515                    position: position + j,
12516                    scale: self.attn_scale,
12517                    eps: self.rms_eps as f32,
12518                    gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12519                    late_qk_norm: self.qk_norm_after_rope,
12520                    output_gate: *output_gate,
12521                    q_norm: q_norm.as_deref(),
12522                    k_norm: k_norm.as_deref(),
12523                    inv_freq: &inv_freq,
12524                    cpu_k: cpu_k.clone(),
12525                    cpu_v: cpu_v.clone(),
12526                    cpu_stored: cpu_stored + j,
12527                    o1: None,
12528                };
12529                if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12530                    return Err(false);
12531                }
12532                if draft_rows < lm.1 {
12533                    if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12534                        return Err(false);
12535                    }
12536                } else {
12537                    graph.encode_lm_head(&m.final_norm, lm);
12538                }
12539                if !graph.encode_argmax(n_arg, j as u32 + 1) {
12540                    return Err(false);
12541                }
12542                if split {
12543                    // CMF_MTP_CHAIN_SPLIT=1: commit every step so the GPU
12544                    // starts on step 0 while the host encodes the rest
12545                    graph.commit();
12546                }
12547            }
12548        }
12549        let t_enc = t_chain.elapsed();
12550        if graph.sync_checked().is_err() {
12551            return Err(true);
12552        }
12553        if std::env::var_os("CMF_GRAPH_SPEC_TIME").is_some() {
12554            eprintln!(
12555                "mtp-chain: encode {:.1} ms | wait {:.1} ms (k={k}, head rows {draft_rows}{})",
12556                t_enc.as_secs_f64() * 1e3,
12557                (t_chain.elapsed() - t_enc).as_secs_f64() * 1e3,
12558                if split { ", split" } else { "" }
12559            );
12560        }
12561        let mut ids = vec![0u32; k];
12562        if !graph.chain_ids_read(&mut ids) {
12563            return Err(true);
12564        }
12565        let mut kbuf = vec![0f32; k * nkv * hd];
12566        let mut vbuf = vec![0f32; k * nkv * hd];
12567        if !crate::gpu_metal::kv_mirror_read_rows(
12568            self.mtp_kv_id(),
12569            Self::MTP_LAYER_BASE,
12570            nkv,
12571            hd,
12572            cpu_stored,
12573            k,
12574            &mut kbuf,
12575            &mut vbuf,
12576        ) {
12577            return Err(true);
12578        }
12579        for r in 0..k {
12580            m.kv.append(
12581                &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12582                &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12583                &[],
12584            );
12585        }
12586        Ok(ids)
12587    }
12588
12589    fn try_batch_graph_wgpu(
12590        &self,
12591        hiddens: &mut [f32],
12592        positions: &[usize],
12593        k: usize,
12594        spec: Option<crate::gpu::SpecTail<'_>>,
12595    ) -> crate::gpu::BatchGraphOutcome {
12596        self.try_batch_graph_wgpu_prefix(hiddens, positions, k, spec, None)
12597    }
12598
12599    /// `try_batch_graph_wgpu` with the device-prefix mode: `layers_run`
12600    /// Some lets a stack that does not fit run its leading layers (the
12601    /// token graph's prefix rule) and reports how many; `hiddens` then
12602    /// holds the boundary rows and the caller runs the rest on the host.
12603    fn try_batch_graph_wgpu_prefix(
12604        &self,
12605        hiddens: &mut [f32],
12606        positions: &[usize],
12607        k: usize,
12608        spec: Option<crate::gpu::SpecTail<'_>>,
12609        layers_run: Option<&mut usize>,
12610    ) -> crate::gpu::BatchGraphOutcome {
12611        let graph_end = match self.mimo_moe.graph_prefix_end() {
12612            Some(end) if end < self.num_layers => {
12613                if layers_run.is_none() || spec.is_some() || end == 0 {
12614                    return crate::gpu::BatchGraphOutcome::Declined;
12615                }
12616                end
12617            }
12618            _ => self.num_layers,
12619        };
12620        let _tb = std::time::Instant::now();
12621        let batch_debug = std::env::var_os("CMF_BATCH_DEBUG").is_some();
12622        if self.attn_softcap > 0.0 {
12623            return crate::gpu::BatchGraphOutcome::Declined; // capped scores: no graph kernel — CPU path
12624        }
12625        // Same attention contract as the token graph: per-layer geometry
12626        // rides `geom`, anything it cannot express declines by name.
12627        if let Some(reason) = self.wgpu_graph_attn_decline() {
12628            self.note_graph_decline("wgpu batch graph", reason);
12629            return crate::gpu::BatchGraphOutcome::Declined;
12630        }
12631        let nh = self.num_heads;
12632        let (nkv, hd, rd) = self.layer_geom(0);
12633        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
12634        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
12635            if let Some((m, i, kind, rs)) = t
12636                .graph_weight()
12637                .or_else(|| t.graph_weight_descriptor())
12638            {
12639                let name = &m.tensors[i].name;
12640                let prism = if crate::prism::is_inverse_embedding(m, name) {
12641                    crate::gpu::GraphPrismOp::InverseEmbedding
12642                } else if crate::prism::is_forward_weight(m, name) {
12643                    crate::gpu::GraphPrismOp::Forward
12644                } else {
12645                    crate::gpu::GraphPrismOp::None
12646                };
12647                return Some(crate::gpu::GraphW {
12648                    idx: i,
12649                    kind,
12650                    row_scale: rs,
12651                    data: &[],
12652                    prism,
12653                    affine: crate::prism::is_affine_target(m, name),
12654                });
12655            }
12656            if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
12657                eprintln!(
12658                    "batch graph: tensor has no graph descriptor/f32 fallback rows={} cols={}",
12659                    t.rows(),
12660                    t.cols()
12661                );
12662            }
12663            t.as_f32().map(|d| crate::gpu::GraphW {
12664                idx: 0,
12665                kind: 4,
12666                row_scale: &[],
12667                data: d,
12668                prism: crate::gpu::GraphPrismOp::None,
12669                affine: false,
12670            })
12671        }
12672        let built: Option<(
12673            Vec<crate::gpu::GraphLayer<'_>>,
12674            std::sync::Arc<cortiq_core::CmfModel>,
12675        )> = (|| {
12676            let mut layers = Vec::with_capacity(graph_end);
12677            let mut model = None;
12678            for li in 0..graph_end {
12679                let lw = &self.weights.layers[self.phys_layer(li)];
12680                // MoE routes per token, so its experts are encoded token by
12681                // token inside the batched submit while attention and the
12682                // projections stay GEMMs. Refusing MoE here is what left
12683                // prefill running one position at a time: 33 tok/s against
12684                // 54 on decode, i.e. reading the prompt was slower than
12685                // writing the answer.
12686                let gffn = match &lw.ffn {
12687                    FfnKind::Dense(d) if !d.segs.is_empty() => {
12688                        if batch_debug {
12689                            eprintln!("batch graph: dense segmented FFN at layer {li}");
12690                        }
12691                        return None;
12692                    }
12693                    FfnKind::Dense(d) => crate::gpu::GraphFfn::Dense {
12694                        gate: gw(&d.gate_proj)?,
12695                        up: gw(&d.up_proj)?,
12696                        down: gw(&d.down_proj)?,
12697                    },
12698                    FfnKind::Moe(m) => {
12699                        // Adaptive τ and expert masks stay on the CPU path.
12700                        // Sigmoid scores, the selection bias, a routed scale
12701                        // ≠ 1 and an ungated shared expert (hy_v3) ride the
12702                        // same flags word as the token graph — before, this
12703                        // refusal sent every Hy-MT2-30B prompt to the chunked
12704                        // fallback (8 tok/s of ingest against 53 of decode).
12705                        if m.route_tau.is_some() || m.mask.is_some() {
12706                            return None;
12707                        }
12708                        // A shared expert rides as slot top_k (gated or
12709                        // not is a flag on the select kernel); without one
12710                        // (MiMo-V2, LFM2-MoE) the kernels run top_k slots.
12711                        let shared = m.shared.as_ref();
12712                        let has_shared = shared.is_some();
12713                        let shared_gated = matches!(shared, Some((_, Some(_))));
12714                        let sgate = match shared {
12715                            Some((_, Some(sg))) => gw(sg)?,
12716                            // Ungated or absent: the router plane stands in
12717                            // so the plumbing stays total; the kernel pins
12718                            // weight 1 or never reads it.
12719                            _ => gw(&m.router)?,
12720                        };
12721                        let router = gw(&m.router)?;
12722                        // The batch MoE kernels still consume raw per-token
12723                        // rows and do not carry the descriptor-aware Prism
12724                        // transform/affine bit for router or shared-gate
12725                        // planes.  Refuse rather than route an untransformed
12726                        // source activation.
12727                        if router.prism != crate::gpu::GraphPrismOp::None
12728                            || router.affine
12729                            || sgate.prism != crate::gpu::GraphPrismOp::None
12730                            || sgate.affine
12731                        {
12732                            return None;
12733                        }
12734                        let inter = m.experts.first()?.gate_proj.rows();
12735                        let mut experts = Vec::with_capacity(m.experts.len() + 1);
12736                        let mut q4tp: Option<bool> = None;
12737                        let mut gu_q2: Option<bool> = None;
12738                        for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
12739                            if !matches!(e.act, Act::Silu)
12740                                || e.gate_proj.rows() != inter
12741                                || e.up_proj.rows() != inter
12742                            {
12743                                return None;
12744                            }
12745                            // Same ladder as the token graph: q4t → q2tp
12746                            // (mixed profile: 2-bit gate/up over a q4tp
12747                            // down) → q4tp. Uniform across the layer.
12748                            let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
12749                                Some((mm, gi)) => (
12750                                    mm,
12751                                    gi,
12752                                    e.up_proj.mapped_q4t()?.1,
12753                                    e.down_proj.mapped_q4t()?.1,
12754                                    false,
12755                                    false,
12756                                ),
12757                                None => match e.gate_proj.mapped_q2tp() {
12758                                    Some((mm, gi)) => (
12759                                        mm,
12760                                        gi,
12761                                        e.up_proj.mapped_q2tp()?.1,
12762                                        e.down_proj.mapped_q4tp()?.1,
12763                                        true,
12764                                        true,
12765                                    ),
12766                                    None => {
12767                                        let (mm, gi) = e.gate_proj.mapped_q4tp()?;
12768                                        (
12769                                            mm,
12770                                            gi,
12771                                            e.up_proj.mapped_q4tp()?.1,
12772                                            e.down_proj.mapped_q4tp()?.1,
12773                                            true,
12774                                            false,
12775                                        )
12776                                    }
12777                                },
12778                            };
12779                            if *q4tp.get_or_insert(is_p) != is_p
12780                                || *gu_q2.get_or_insert(is_q2) != is_q2
12781                            {
12782                                return None;
12783                            }
12784                            if [gi, ui, di].into_iter().any(|idx| {
12785                                mm.tensors
12786                                    .get(idx)
12787                                    .is_some_and(|t| {
12788                                        crate::prism::is_forward_weight(mm, &t.name)
12789                                            || crate::prism::is_affine_target(mm, &t.name)
12790                                    })
12791                            }) {
12792                                return None;
12793                            }
12794                            model.get_or_insert_with(|| mm.clone());
12795                            experts.push((gi, ui, di));
12796                        }
12797                        crate::gpu::GraphFfn::Moe {
12798                            router,
12799                            shared_gate: sgate,
12800                            experts,
12801                            n_exp: m.experts.len(),
12802                            top_k: m.top_k,
12803                            inter,
12804                            norm_topk: m.norm_topk_prob,
12805                            q4tp: q4tp?,
12806                            gu_q2: gu_q2.unwrap_or(false),
12807                            sigmoid: m.router_sigmoid,
12808                            bias: m.expert_bias.as_deref(),
12809                            has_shared,
12810                            shared_gated,
12811                            route_scale: m.routed_scaling,
12812                        }
12813                    }
12814                    _ => return None,
12815                };
12816                let attn = match &lw.attn {
12817                    AttnKind::Full {
12818                        wq,
12819                        wk,
12820                        wv,
12821                        wo,
12822                        q_norm,
12823                        k_norm,
12824                        output_gate,
12825                        softplus_gate,
12826                        bias,
12827                    } => {
12828                        if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
12829                            if batch_debug {
12830                                eprintln!(
12831                                    "batch graph: unsupported Full attention gate at layer {li} softplus={} heads={}",
12832                                    softplus_gate.is_some(),
12833                                    self.attention_heads_per_layer.is_some()
12834                                );
12835                            }
12836                            return None;
12837                        }
12838                        let (m, _, _, _) = wq
12839                            .graph_weight()
12840                            .or_else(|| wq.graph_weight_descriptor())?;
12841                        model = Some(m.clone());
12842                        crate::gpu::GraphAttn::Full {
12843                            wq: gw(wq)?,
12844                            wk: gw(wk)?,
12845                            wv: gw(wv)?,
12846                            wo: gw(wo)?,
12847                            q_norm: q_norm.as_deref(),
12848                            k_norm: k_norm.as_deref(),
12849                            late_qk_norm: self.qk_norm_after_rope,
12850                            bias: bias
12851                                .as_ref()
12852                                .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
12853                            output_gate: *output_gate,
12854                            cpu_k: self.kv_cache.layers[li].k_heads(),
12855                            cpu_v: self.kv_cache.layers[li].v_heads(),
12856                            geom: self.graph_attn_geom(li),
12857                        }
12858                    }
12859                    AttnKind::LinearGdn(w) => {
12860                        let Some(cfg) = self.gdn_cfg else {
12861                            if batch_debug {
12862                                eprintln!("batch graph: no GDN config at layer {li}");
12863                            }
12864                            return None;
12865                        };
12866                        let (m, _, _, _) = w
12867                            .in_proj_qkv
12868                            .graph_weight()
12869                            .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
12870                        model = Some(m.clone());
12871                        crate::gpu::GraphAttn::Gdn {
12872                            qkv: gw(&w.in_proj_qkv)?,
12873                            z: gw(&w.in_proj_z)?,
12874                            a: gw(&w.in_proj_a)?,
12875                            b: gw(&w.in_proj_b)?,
12876                            out: gw(&w.out_proj)?,
12877                            conv1d: &w.conv1d,
12878                            a_log: &w.a_log,
12879                            dt_bias: &w.dt_bias,
12880                            norm: &w.norm,
12881                            nv: cfg.num_v_heads,
12882                            nk: cfg.num_k_heads,
12883                            dk: cfg.key_head_dim,
12884                            dv: cfg.value_head_dim,
12885                            kk: cfg.conv_kernel,
12886                            cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
12887                        }
12888                    }
12889                    _ => return None,
12890                };
12891                layers.push(crate::gpu::GraphLayer {
12892                    input_norm: &lw.input_norm,
12893                    attn,
12894                    post_norm: &lw.post_norm,
12895                    ffn: gffn,
12896                });
12897            }
12898            Some((layers, model?))
12899        })();
12900        let Some((layers, model)) = built else {
12901            {
12902                use std::sync::atomic::{AtomicBool, Ordering};
12903                static SAID: AtomicBool = AtomicBool::new(false);
12904                if !SAID.swap(true, Ordering::Relaxed) {
12905                    tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
12906                }
12907            }
12908            return crate::gpu::BatchGraphOutcome::Declined;
12909        };
12910        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
12911            eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
12912        }
12913        crate::gpu::forward_batch_graph(
12914            &model,
12915            self.graph_kv_id,
12916            &layers,
12917            &self.inv_freq,
12918            hiddens,
12919            nh,
12920            nkv,
12921            hd,
12922            rd,
12923            self.hidden_size,
12924            self.intermediate_size,
12925            positions,
12926            self.kv_cache.max_seq_len,
12927            gemma,
12928            self.rms_eps as f32,
12929            self.attn_scale,
12930            k,
12931            &(0..graph_end)
12932                .map(|li| self.kv_cache.layers[self.phys_layer(li)].o1_views())
12933                .collect::<Vec<_>>(),
12934            self.o1_epoch,
12935            spec,
12936            layers_run,
12937        )
12938    }
12939
12940    /// Same, stopping after layer `upto` inclusive (routing probe φ).
12941    /// `CMF_DSV4_DRAFT_PROBE=1` — grade the draft against what the trunk goes on
12942    /// to produce. Off by default; it runs a whole draft per decoded token.
12943    fn draft_probe() -> bool {
12944        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12945        *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
12946    }
12947
12948    /// `CMF_DSV4_DRAFT_PROBE=1`: measure how much of the draft the trunk
12949    /// would have agreed with, WITHOUT verifying or rolling anything back.
12950    ///
12951    /// The number this produces decides the whole speculation design — at
12952    /// acceptance a, a block of B positions yields 1 + a + a² + ... tokens
12953    /// per trunk pass — so it is worth measuring before any of the machinery
12954    /// that would exploit it exists. Each draft is parked with the position
12955    /// it was made at, and graded as the real tokens arrive.
12956    /// `CMF_DSV4_SPEC=1` — the DeepSeek-V4 speculative decode: draft five
12957    /// on the card, verify them in one batched trunk pass, commit the
12958    /// accepted prefix, roll the rest back.
12959    #[cfg(feature = "gpu")]
12960    fn dsv4_spec_on() -> bool {
12961        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12962        *ON.get_or_init(|| {
12963            // Test-only runtime gate: model loading still performs the same
12964            // reservation and trunk packing, which gives rollback parity a
12965            // topology-identical non-speculative control arm.
12966            if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
12967                return v != "0";
12968            }
12969            // An explicit value is a diagnostic force/escape hatch.  With no
12970            // knob, speculation is eligible only when model loading reserved
12971            // its bounded pack.  On small q4tp cards the geometric reserve
12972            // gate deliberately leaves this at zero: trying to build DSpark
12973            // after the exact trunk filled VRAM is both slower and a device
12974            // OOM (measured on A40).
12975            std::env::var("CMF_DSV4_SPEC")
12976                .map(|v| v != "0")
12977                .unwrap_or_else(|_| {
12978                    crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
12979                })
12980        })
12981    }
12982
12983    /// One speculative round at the decode tip. `t_next` is the token the
12984    /// sampler just committed for `next_pos`. Returns the EXTRA accepted
12985    /// tokens (possibly none) and the new position, with `graph_logits`
12986    /// left holding the last accepted position's logits — exactly what the
12987    /// loop top expects. `None` means "speculate not this round": nothing
12988    /// was committed, the caller forwards normally.
12989    #[cfg(feature = "gpu")]
12990    fn dsv4_spec_step(
12991        &mut self,
12992        tip_token: u32,
12993        t_next: u32,
12994        next_pos: usize,
12995        max_extra: usize,
12996        drafted: &mut usize,
12997        accepted_ctr: &mut usize,
12998    ) -> Option<(Vec<u32>, usize)> {
12999        let t_all = std::time::Instant::now();
13000        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13001            thread_local! {
13002                static LAST: std::cell::Cell<Option<std::time::Instant>> =
13003                    const { std::cell::Cell::new(None) };
13004            }
13005            LAST.with(|l| {
13006                if let Some(prev) = l.get() {
13007                    eprintln!(
13008                        "между раундами {:.1} мс",
13009                        prev.elapsed().as_secs_f64() * 1e3
13010                    );
13011                }
13012                l.set(Some(std::time::Instant::now()));
13013            });
13014        }
13015        if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13016            eprintln!("spec_step: вход pos={next_pos}");
13017        }
13018        let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
13019        let cfg = self.dsv4.as_ref().map(|b| b.2)?;
13020        // The draft state and its capture, armed exactly as the probe does.
13021        if self.dspark.is_none() {
13022            let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13023            if t.is_empty() {
13024                return None;
13025            }
13026            crate::dsv4::dspark_arm(&t, cfg.dim);
13027            self.dspark = Some(crate::dsv4::DsparkState::new(
13028                self.dsv4_mtp.len(),
13029                &cfg,
13030                t.len(),
13031            ));
13032        }
13033        let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13034        let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
13035        if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13036            eprintln!("spec_step: пак не построился (targets {targets:?})");
13037        }
13038        let pack = pack?;
13039        let block = crate::dsv4::dspark_block();
13040        let b_box = self.dsv4.as_mut()?;
13041        let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
13042        let ds = self.dspark.as_mut()?;
13043        // The tip's captures: either this token ran on a normal path that
13044        // filled the thread-local, or the previous spec round left them.
13045        let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
13046        if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
13047            if dbg {
13048                eprintln!("spec_step: нет захвата");
13049            }
13050            return None;
13051        }
13052        ds.have_hidden = true;
13053        let tip_pos = next_pos.checked_sub(1)?;
13054        let draft_started = std::time::Instant::now();
13055        let mut conf = Vec::new();
13056        let props = crate::dsv4::dspark_draft_gpu(
13057            g,
13058            &self.dsv4_mtp,
13059            &cfg,
13060            ds,
13061            pack,
13062            st.kv_id,
13063            tip_token,
13064            tip_pos,
13065            self.pool.as_deref(),
13066            &mut conf,
13067        );
13068        self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13069        *drafted += block;
13070        if props.is_empty() || props[0] != t_next {
13071            if dbg {
13072                eprintln!(
13073                    "spec_step: черновик {} (props0={:?} t_next={t_next})",
13074                    if props.is_empty() {
13075                        "пуст"
13076                    } else {
13077                        "мимо"
13078                    },
13079                    props.first()
13080                );
13081            }
13082            return None;
13083        }
13084        // `fed[0]` is `t_next`, which the outer loop has already committed;
13085        // only `fed[1..]` become additional output tokens. Cap the verify
13086        // transaction itself to the caller's remaining output budget instead
13087        // of merely truncating the returned vector: otherwise the KV/state
13088        // would advance past `max_tokens` and a 64-token request could return
13089        // 66 tokens (and poison a reused session with two invisible steps).
13090        let mut k_verify = crate::dsv4::dspark_verify_k()
13091            .min(props.len())
13092            .min(max_extra.saturating_add(1));
13093        // Adaptive depth: positions the draft itself doubts are paid for on
13094        // every verify and delivered almost never (natural-text survival
13095        // [.67 .50 .29 .08 .04]). `CMF_DSPARK_CONF_MIN=p` trims the fed
13096        // prefix at the first proposal whose confidence drops below p; on
13097        // predictable text the confidences stay high and nothing changes.
13098        let conf_min = {
13099            static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
13100            *M.get_or_init(|| {
13101                std::env::var("CMF_DSPARK_CONF_MIN")
13102                    .ok()
13103                    .and_then(|v| v.parse().ok())
13104                    .unwrap_or(0.0)
13105            })
13106        };
13107        if conf_min > 0.0 && conf.len() >= props.len() {
13108            let mut keep = 1usize;
13109            while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
13110                keep += 1;
13111            }
13112            k_verify = k_verify.min(keep.max(2));
13113        }
13114        if k_verify < 2 {
13115            return None;
13116        }
13117        let mut fed = Vec::with_capacity(k_verify);
13118        fed.push(t_next);
13119        fed.extend_from_slice(&props[1..k_verify]);
13120        let mut argmax = Vec::new();
13121        let mut logits_all = Vec::new();
13122        let mut walked = Vec::new();
13123        let txn = crate::dsv4::dsv4_verify_chunk(
13124            g,
13125            layers,
13126            &cfg,
13127            st,
13128            &fed,
13129            next_pos,
13130            &self.inv_freq,
13131            self.pool.as_deref(),
13132            &targets,
13133            &mut argmax,
13134            &mut logits_all,
13135            &mut walked,
13136        );
13137        if txn.is_none() && dbg {
13138            eprintln!("spec_step: verify отказал");
13139        }
13140        let txn = txn?;
13141        let spec_gpu_end = txn.gpu_end;
13142        let b = fed.len();
13143        let mut accepted = 1usize;
13144        while accepted < b && fed[accepted] == argmax[accepted - 1] {
13145            accepted += 1;
13146        }
13147        // `CMF_DSV4_SPEC_FORCE_REJECT=1` — accept nothing beyond the known
13148        // token, every round: the pure rollback exerciser. The output must
13149        // stay byte-identical to the plain walk; anything else is a
13150        // transaction bug, isolated from the acceptance logic.
13151        if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
13152            accepted = 1;
13153        }
13154        if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
13155            eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
13156        }
13157        let t_fin = std::time::Instant::now();
13158        if !crate::dsv4::dsv4_spec_finish(
13159            g,
13160            layers,
13161            &cfg,
13162            st,
13163            txn,
13164            accepted,
13165            &fed,
13166            &self.inv_freq,
13167            self.pool.as_deref(),
13168        ) {
13169            tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
13170            return None;
13171        }
13172        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13173            eprintln!(
13174                "finish(k={accepted}): {:.1} мс",
13175                t_fin.elapsed().as_secs_f64() * 1e3
13176            );
13177        }
13178        *accepted_ctr += accepted - 1;
13179        // Captures per accepted token: device targets photographed by the
13180        // batch, host targets from the verify's own walk. The last one
13181        // becomes the new tip's draft input; every one owes the ring an
13182        // entry for its position.
13183        let (hc, dim) = (cfg.hc_mult, cfg.dim);
13184        // Complete-chain layers are photographed by the fused submission;
13185        // partial device layers overwrite that slot after exact host cold-
13186        // expert correction.  Thus every target in the contiguous device
13187        // prefix has a valid per-token capture.
13188        let dev_caps: Vec<usize> = targets
13189            .iter()
13190            .copied()
13191            .filter(|&t| t < spec_gpu_end)
13192            .collect();
13193        let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
13194        if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
13195            return None;
13196        }
13197        for t in 0..accepted {
13198            let tip = t + 1 == accepted;
13199            for (slot, &tl) in targets.iter().enumerate() {
13200                if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
13201                    let lo = (di * b + t) * hc * dim;
13202                    crate::dsv4::dspark_capture(
13203                        &caps_all[lo..lo + hc * dim],
13204                        &cfg,
13205                        slot,
13206                        &mut ds.main_hidden,
13207                    );
13208                } else if tip
13209                    && crate::dsv4::dspark_peek_slot(slot, dim, {
13210                        let lo = slot * dim;
13211                        &mut ds.main_hidden[lo..lo + dim]
13212                    })
13213                {
13214                    // The tip's host-layer captures are the walk's own
13215                    // per-layer notes — exact. (The walk that ran last ended
13216                    // on exactly this token, on both the accept-all and the
13217                    // rollback path.)
13218                } else {
13219                    // Intermediate tokens: the post-tail state stands in for
13220                    // the per-layer capture on host targets below the last
13221                    // layer. Ring-entry quality only; the tip is exact.
13222                    crate::dsv4::dspark_capture(
13223                        &walked[t * hc * dim..(t + 1) * hc * dim],
13224                        &cfg,
13225                        slot,
13226                        &mut ds.main_hidden,
13227                    );
13228                }
13229            }
13230            crate::dsv4::dspark_ring_append(
13231                g,
13232                &self.dsv4_mtp,
13233                &cfg,
13234                ds,
13235                next_pos + t,
13236                self.pool.as_deref(),
13237            );
13238        }
13239        let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
13240        self.graph_logits = Some(row);
13241        // The speculative loop never runs the probe, so the trunk tally has
13242        // no other place to cycle. Armed only when someone asked for the
13243        // dump; the host tail is the only tallying path here, which is
13244        // precisely the population a partial pack would serve.
13245        if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
13246            crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
13247            crate::dsv4::pick_tally_arm();
13248        }
13249        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13250            eprintln!(
13251                "spec_step total {:.1} мс (k={accepted})",
13252                t_all.elapsed().as_secs_f64() * 1e3
13253            );
13254        }
13255        Some((fed[1..accepted].to_vec(), next_pos + accepted))
13256    }
13257
13258    fn dspark_probe(&mut self, position: usize, token_id: u32) {
13259        if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
13260            return;
13261        }
13262        // What the trunk just routed to, for this token.
13263        let trunk_now = crate::dsv4::pick_tally_take();
13264        crate::dsv4::trunk_freq_note(&trunk_now);
13265        if !trunk_now.is_empty() {
13266            self.dspark_trunk_picks.push(trunk_now);
13267            let keep = crate::dsv4::dspark_block();
13268            if self.dspark_trunk_picks.len() > keep {
13269                self.dspark_trunk_picks.remove(0);
13270            }
13271        }
13272        // Grade whatever is waiting: the token just decoded sits at
13273        // `position`, so it answers the draft made at `position - 1 - i`.
13274        for p in std::mem::take(&mut self.dspark_pending) {
13275            let Some(i) = position.checked_sub(p.0 + 1) else {
13276                continue;
13277            };
13278            let mut p = p;
13279            if i < p.1.len() {
13280                if p.2 && p.1[i] == token_id {
13281                    p.3 = i + 1;
13282                } else {
13283                    p.2 = false;
13284                }
13285                if i + 1 < p.1.len() {
13286                    self.dspark_pending.push(p);
13287                    continue;
13288                }
13289            }
13290            self.dspark_hist.push(p.3);
13291            self.dspark_real.push(token_id);
13292        }
13293        let Some(b) = &mut self.dsv4 else { return };
13294        let (g, layers, cfg) = (&b.0, &b.1, b.2);
13295        let n_layers = layers.len();
13296        if self.dspark.is_none() {
13297            let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13298            if t.is_empty() {
13299                return;
13300            }
13301            eprintln!(
13302                "DSpark: захват со слоёв {t:?}, блок {}",
13303                crate::dsv4::dspark_block()
13304            );
13305            crate::dsv4::dspark_arm(&t, cfg.dim);
13306            self.dspark = Some(crate::dsv4::DsparkState::new(
13307                self.dsv4_mtp.len(),
13308                &cfg,
13309                t.len(),
13310            ));
13311        }
13312        let ds = self.dspark.as_mut().unwrap();
13313        if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
13314            return; // this token ran on a path that captures nothing
13315        }
13316        let mut conf = Vec::new();
13317        crate::dsv4::pick_tally_arm();
13318        // The trunk has already consumed the adaptive VRAM budget. Until the
13319        // draft owns an explicit bounded device pack, its tensors are an
13320        // out-of-core CPU/disk tier by contract: never let per-op probes try
13321        // to squeeze another multi-gigabyte MTP expert cache onto the card.
13322        let draft_started = std::time::Instant::now();
13323        #[cfg(feature = "gpu")]
13324        let gpu_draft = crate::dsv4::dspark_gpu_on();
13325        #[cfg(not(feature = "gpu"))]
13326        let gpu_draft = false;
13327        let props = if gpu_draft {
13328            #[cfg(feature = "gpu")]
13329            {
13330                let kv_id = b.3.kv_id;
13331                match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
13332                    Some(pk) => crate::dsv4::dspark_draft_gpu(
13333                        g,
13334                        &self.dsv4_mtp,
13335                        &cfg,
13336                        ds,
13337                        pk,
13338                        kv_id,
13339                        token_id,
13340                        position,
13341                        self.pool.as_deref(),
13342                        &mut conf,
13343                    ),
13344                    None => Vec::new(),
13345                }
13346            }
13347            #[cfg(not(feature = "gpu"))]
13348            Vec::new()
13349        } else {
13350            crate::gpu::cpu_scope(|| {
13351                crate::dsv4::dspark_draft(
13352                    g,
13353                    &self.dsv4_mtp,
13354                    &cfg,
13355                    ds,
13356                    token_id,
13357                    position,
13358                    self.pool.as_deref(),
13359                    &mut conf,
13360                )
13361            })
13362        };
13363        self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13364        let draft_picks = crate::dsv4::pick_tally_take();
13365        crate::dsv4::dspark_freq_note(&draft_picks);
13366        // Re-arm for the NEXT trunk token; the probe runs after the forward,
13367        // so this is the only place that can.
13368        crate::dsv4::pick_tally_arm();
13369        if !props.is_empty() {
13370            // Two ratios, side by side: what a batched verify over the trunk
13371            // would read against what it asks for, and the same for the
13372            // draft's three stages. Near 1.0 means a batch amortises nothing.
13373            let (tu, tt) = {
13374                let flat: Vec<(usize, Vec<usize>)> = self
13375                    .dspark_trunk_picks
13376                    .iter()
13377                    .flat_map(|v| v.iter().cloned())
13378                    .collect();
13379                // Per layer, across the window of tokens.
13380                let mut per: std::collections::HashMap<usize, Vec<usize>> =
13381                    std::collections::HashMap::new();
13382                for (li, picks) in flat {
13383                    per.entry(li).or_default().extend(picks);
13384                }
13385                let n = per.len().max(1);
13386                let mut u = 0usize;
13387                let mut t = 0usize;
13388                for (_, v) in per {
13389                    t += v.len();
13390                    u += v.iter().collect::<std::collections::HashSet<_>>().len();
13391                }
13392                (u / n, t / n)
13393            };
13394            let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
13395            self.dspark_exp.push((tu, tt, du, dt));
13396            self.dspark_pending.push((position, props, true, 0));
13397        }
13398        if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
13399            let n = self.dspark_hist.len() as f32;
13400            let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
13401            let block = crate::dsv4::dspark_block();
13402            let mut at = vec![0usize; block + 1];
13403            for &k in &self.dspark_hist {
13404                at[k] += 1;
13405            }
13406            // Prefix survival: S_i = P(the first i positions all held).
13407            let mut surv = Vec::with_capacity(block);
13408            for i in 1..=block {
13409                let k = at[i..].iter().sum::<usize>() as f32 / n;
13410                surv.push(format!("{k:.2}"));
13411            }
13412            let distinct = self
13413                .dspark_real
13414                .iter()
13415                .collect::<std::collections::HashSet<_>>()
13416                .len();
13417            let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
13418                (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
13419            });
13420            let m = self.dspark_exp.len().max(1);
13421            eprintln!(
13422                "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
13423                 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
13424                self.dspark_hist.len(),
13425                mean + 1.0,
13426                surv.join(" ")
13427            );
13428            eprintln!(
13429                "DSpark: разных токенов {distinct} из {} (вырожденность), \
13430                 эксперты ствол {}/{} на слой за {block} токенов, \
13431                 черновик {}/{} за блок, draft {:.2} мс/блок",
13432                self.dspark_real.len(),
13433                tu / m,
13434                tt / m,
13435                du / m,
13436                dt / m,
13437                self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
13438            );
13439        }
13440    }
13441
13442    fn forward_layers_upto(
13443        &mut self,
13444        hidden: &[f32],
13445        position: usize,
13446        task_mask: Option<&TaskMask>,
13447        upto: Option<usize>,
13448    ) -> Vec<f32> {
13449        // In-process multi-GPU: each segment runs pinned to its card,
13450        // and the only thing crossing the boundary is one hidden vector
13451        // that never leaves this address space. Same layer split the
13452        // network mode does, minus the second process, the socket, the
13453        // serialization and the dir_hash handshake.
13454        if let Some(plan) = self.gpu_plan.clone() {
13455            if upto.is_none() && plan.len() > 1 {
13456                let mut h = hidden.to_vec();
13457                for &(dev, from, upto_incl) in plan.iter() {
13458                    h = crate::gpu::with_device(dev, || {
13459                        self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
13460                    });
13461                }
13462                return h;
13463            }
13464        }
13465        self.forward_layers_span(hidden, position, task_mask, 0, upto)
13466    }
13467
13468    /// Split this pipeline's layer stack across local GPUs: segment i
13469    /// runs on `devices[i]`. Contiguous and even by layer count — the
13470    /// VRAM-weighted planner is the next step, and an uneven card pair
13471    /// is why it will be needed. `None` clears the plan.
13472    pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
13473        self.set_gpu_plan_at(devices, None)
13474    }
13475
13476    /// The same, with an explicit first boundary (`--peer-split`): card
13477    /// 0 takes layers `[0..at)`, the rest split what remains. Uneven
13478    /// cards, or an attention-heavy head, are why this knob exists.
13479    pub fn set_gpu_plan_at(
13480        &mut self,
13481        devices: Option<&[usize]>,
13482        at: Option<usize>,
13483    ) -> Result<(), String> {
13484        let Some(devs) = devices.filter(|d| d.len() > 1) else {
13485            self.gpu_plan = None;
13486            return Ok(());
13487        };
13488        self.split_supported()?;
13489        let n = self.num_layers;
13490        if devs.len() > n {
13491            return Err(format!("{} devices for {n} layers", devs.len()));
13492        }
13493        if let Some(k) = at {
13494            if k == 0 || k >= n {
13495                return Err(format!("split at {k}: the model has {n} layers"));
13496            }
13497            if devs.len() == 2 {
13498                self.gpu_plan = Some(std::sync::Arc::new(vec![
13499                    (devs[0], 0, k - 1),
13500                    (devs[1], k, n - 1),
13501                ]));
13502                return Ok(());
13503            }
13504            return Err(format!(
13505                "an explicit split point takes exactly 2 devices, got {}",
13506                devs.len()
13507            ));
13508        }
13509        let per = n.div_ceil(devs.len());
13510        let mut plan = Vec::with_capacity(devs.len());
13511        let mut from = 0usize;
13512        for &d in devs {
13513            if from >= n {
13514                break;
13515            }
13516            let upto = (from + per - 1).min(n - 1);
13517            plan.push((d, from, upto));
13518            from = upto + 1;
13519        }
13520        self.gpu_plan = Some(std::sync::Arc::new(plan));
13521        Ok(())
13522    }
13523
13524    /// The active in-process split, if any: (device, first layer, last).
13525    pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
13526        self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
13527    }
13528
13529    /// Layer span [from ..= upto] (upto None = last layer): the building
13530    /// block the network pipeline-split rides on. `from > 0` skips the
13531    /// arch escape hatches (the pub `forward_span` refuses those archs
13532    /// first) and the whole-token graph — the plain per-layer loop is
13533    /// the canonical executor for a partial stack.
13534    fn embryo_qtensor_f32(t: &QTensor) -> Vec<f32> {
13535        if let Some(x) = t.as_f32() {
13536            return x.to_vec();
13537        }
13538        let mut out = vec![0.0; t.rows() * t.cols()];
13539        for r in 0..t.rows() {
13540            t.row_f32(r, &mut out[r * t.cols()..(r + 1) * t.cols()]);
13541        }
13542        out
13543    }
13544
13545    fn embryo_resident_eligible(&self) -> bool {
13546        // One mixer family per file: vmf_phase (kind 0/1) or
13547        // gated_delta_net (kind 4); the anchors are full (2) or bounded (3).
13548        if (self.vmf_cfg.is_none() && self.gdn_cfg.is_none())
13549            || self.num_layers != self.physical_layers
13550            || self.loop_final_norm
13551            || self.weights.layers.len() != self.num_layers
13552            || self.head_clusters.is_none()
13553            || self.final_softcap.is_some()
13554            || self.logit_multiplier.is_some()
13555            || self.attn_softcap != 0.0
13556            || self.mtp.is_some()
13557            || self.g3n.is_some()
13558            || self.dsv4.is_some()
13559            || self.dsv41.is_some()
13560            || self.qwen4_exp.is_some()
13561            // Dynamic routing swaps FFN weights mid-sequence under the
13562            // packed graph; a blend has no single overlay. A STATIC skill
13563            // (`from_model_with_skill`) is fine: the pack reads the live
13564            // `weights.layers[*].ffn`, i.e. the skill's tensors, and a
13565            // later `set_active_skill` change drops the pack
13566            // (`invalidate_for_weight_change`).
13567            || self.dyn_router.is_some()
13568            || self.dyn_phi_layer.is_some()
13569            || self.dyn_blend_loaded
13570            || self.o1_cfg.is_some()
13571            || self.swa.is_some()
13572            || self.sliding_layers.is_some()
13573            || self.global_attn.is_some()
13574            || self.attention_heads_per_layer.is_some()
13575            || self.attn_v_norm
13576            || self
13577                .kv_cache
13578                .layers
13579                .iter()
13580                .any(|l| l.mode != crate::kv_cache::KvMode::F32)
13581            || self.rope_scale != 1.0
13582            || self.rope_scale_local != 1.0
13583            || self.attn_scale != 1.0 / (self.head_dim as f32).sqrt()
13584            || self.hidden_size == 0
13585            || self.hidden_size > 1024
13586            || self.intermediate_size > 1024
13587            || self.num_heads == 0
13588            || self.num_kv_heads == 0
13589            || self.num_heads % self.num_kv_heads != 0
13590            || self.num_heads.saturating_mul(self.head_dim) > 1024
13591            || self.num_kv_heads.saturating_mul(self.head_dim) > 1024
13592            || self.vocab_size == 0
13593            || self.kv_cache.max_seq_len == 0
13594            || self.rotary_dim == 0
13595            || self.rotary_dim > self.head_dim
13596            || self.rotary_dim % 2 != 0
13597            || self.inv_freq.len() < self.rotary_dim / 2
13598        {
13599            return false;
13600        }
13601        // The resident shader is deliberately an f32 profile.  Dequantizing
13602        // a Q4/Q8 tensor into the packed buffer would silently change the
13603        // operator relative to the CPU quantized path, so quantized CMFs
13604        // retain the exact ordinary executor instead of claiming parity.
13605        // Measured consequence (RTX PRO 4000, S4 bounded export requantized
13606        // with `cortiq requant --quant q4tp-quantize`): `eligible=false`,
13607        // the generic wgpu whole-token graph refuses too, and the per-op
13608        // path decodes at ~73 tok/s against ~200 tok/s on the CPU q4tp
13609        // path — a q4tp Embryo-O1 file is a CPU artifact today; the
13610        // resident graph serves the f32 export.
13611        if self.weights.lm_head.as_f32().is_none()
13612            || self.weights.embed_tokens.as_f32().is_none()
13613            || self.weights.lm_head.rows() < self.vocab_size
13614            || self.weights.lm_head.cols() != self.hidden_size
13615            || self.weights.embed_tokens.rows() < self.vocab_size
13616            || self.weights.embed_tokens.cols() != self.hidden_size
13617            || self.weights.final_norm.len() != self.hidden_size
13618        {
13619            return false;
13620        }
13621        if let Some(cfg) = self.vmf_cfg {
13622            if cfg.num_heads.saturating_mul(cfg.nphase) > 1024
13623                || cfg.num_heads.saturating_mul(cfg.nphase.saturating_add(1)) > 1024
13624                || cfg.num_heads.saturating_mul(cfg.value_head_dim) > 1024
13625                || cfg.state_len() == 0
13626            {
13627                return false;
13628            }
13629        }
13630        if let Some(g) = self.gdn_cfg {
13631            // The resident GDN kernels (gpu_wgpu.rs `embryo_core_gdn_*`):
13632            // fused projection ≤ 2048 rows, nv·dv ≤ 1024, dk ≤ 128 lanes,
13633            // dv ≤ 256 lanes in vec4 rows, SiLU output gate (the Embryo
13634            // export), same rms eps as the stack.
13635            if g.num_v_heads == 0
13636                || g.num_k_heads == 0
13637                || g.num_v_heads % g.num_k_heads != 0
13638                || g.key_head_dim == 0
13639                || g.key_head_dim > 128
13640                || g.value_head_dim == 0
13641                || g.value_head_dim > 256
13642                || g.value_head_dim % 4 != 0
13643                || g.conv_kernel == 0
13644                || g.num_v_heads > 512
13645                || g.num_v_heads.saturating_mul(g.value_head_dim) > 1024
13646                || g.conv_dim() > 2048
13647                || g.conv_dim() % 4 != 0
13648                || g.hidden_size != self.hidden_size
13649                || g.output_gate_sigmoid
13650                || g.rms_eps != self.rms_eps
13651                || g.state_len() == 0
13652            {
13653                return false;
13654            }
13655        }
13656        let mut full_seen = false;
13657        for lw in &self.weights.layers {
13658            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
13659                return false;
13660            }
13661            match &lw.attn {
13662                AttnKind::LinearGdn(w) => {
13663                    let Some(g) = self.gdn_cfg else {
13664                        return false;
13665                    };
13666                    let (nv, dv, kk) = (g.num_v_heads, g.value_head_dim, g.conv_kernel);
13667                    if w.in_proj_qkv.rows() != g.conv_dim()
13668                        || w.in_proj_qkv.cols() != self.hidden_size
13669                        || w.in_proj_qkv.as_f32().is_none()
13670                        || w.in_proj_z.rows() != nv * dv
13671                        || w.in_proj_z.cols() != self.hidden_size
13672                        || w.in_proj_z.as_f32().is_none()
13673                        || w.in_proj_a.rows() != nv
13674                        || w.in_proj_a.cols() != self.hidden_size
13675                        || w.in_proj_a.as_f32().is_none()
13676                        || w.in_proj_b.rows() != nv
13677                        || w.in_proj_b.cols() != self.hidden_size
13678                        || w.in_proj_b.as_f32().is_none()
13679                        || w.conv1d.len() != g.conv_dim() * kk
13680                        || w.a_log.len() != nv
13681                        || w.dt_bias.len() != nv
13682                        || w.norm.len() != dv
13683                        || w.out_proj.rows() != self.hidden_size
13684                        || w.out_proj.cols() != nv * dv
13685                        || w.out_proj.as_f32().is_none()
13686                    {
13687                        return false;
13688                    }
13689                }
13690                AttnKind::Linear(w) => {
13691                    let Some(cfg) = self.vmf_cfg else {
13692                        return false;
13693                    };
13694                    if w.thq.rows() != cfg.num_heads * cfg.nphase
13695                        || w.thq.cols() != self.hidden_size
13696                        || w.thq.as_f32().is_none()
13697                        || w.thk.rows() != cfg.num_heads * cfg.nphase
13698                        || w.thk.cols() != self.hidden_size
13699                        || w.thk.as_f32().is_none()
13700                        || w.v_proj.rows() != cfg.num_heads * cfg.value_head_dim
13701                        || w.v_proj.cols() != self.hidden_size
13702                        || w.v_proj.as_f32().is_none()
13703                        || w.out_proj.rows() != self.hidden_size
13704                        || w.out_proj.cols() != cfg.num_heads * cfg.value_head_dim
13705                        || w.out_proj.as_f32().is_none()
13706                        || w.decay.len() != cfg.num_heads * 2 * cfg.nphase
13707                    {
13708                        return false;
13709                    }
13710                    if let Some((kg, kb)) = &w.k_gate {
13711                        if kg.rows() != cfg.num_heads
13712                            || kg.cols() != self.hidden_size
13713                            || kg.as_f32().is_none()
13714                            || kb.len() != cfg.num_heads
13715                        {
13716                            return false;
13717                        }
13718                    }
13719                    if let Some(conv) = &w.conv {
13720                        if conv.len() % self.hidden_size != 0 || conv.len() / self.hidden_size < 2 {
13721                            return false;
13722                        }
13723                    }
13724                }
13725                AttnKind::Full {
13726                    wq,
13727                    wk,
13728                    wv,
13729                    wo,
13730                    q_norm,
13731                    k_norm,
13732                    output_gate,
13733                    softplus_gate,
13734                    bias,
13735                } => {
13736                    if full_seen
13737                        || q_norm.is_some()
13738                        || k_norm.is_some()
13739                        || *output_gate
13740                        || softplus_gate.is_some()
13741                        || bias.is_some()
13742                        || wq.as_f32().is_none()
13743                        || wk.as_f32().is_none()
13744                        || wv.as_f32().is_none()
13745                        || wo.as_f32().is_none()
13746                        || wq.rows() != self.num_heads * self.head_dim
13747                        || wk.rows() != self.num_kv_heads * self.head_dim
13748                        || wv.rows() != self.num_kv_heads * self.head_dim
13749                        || wq.cols() != self.hidden_size
13750                        || wk.cols() != self.hidden_size
13751                        || wv.cols() != self.hidden_size
13752                        || wo.rows() != self.hidden_size
13753                        || wo.cols() != self.num_heads * self.head_dim
13754                    {
13755                        return false;
13756                    }
13757                    full_seen = true;
13758                }
13759                AttnKind::Bounded(w) => {
13760                    // The resident bounded attend scores S + W lanes in one
13761                    // 256-lane chunk; the format caps S + W at 160.
13762                    let Some(ac) = self.anchor_core.as_ref() else {
13763                        return false;
13764                    };
13765                    if self.bounded_rope.is_none()
13766                        || w.window != ac.window
13767                        || w.sink != ac.sink
13768                        || w.window == 0
13769                        || w.window + w.sink > 256
13770                        || w.sink_k.len() != self.num_kv_heads * w.sink * self.head_dim
13771                        || w.sink_v.len() != self.num_kv_heads * w.sink * self.head_dim
13772                        || w.wq.as_f32().is_none()
13773                        || w.wk.as_f32().is_none()
13774                        || w.wv.as_f32().is_none()
13775                        || w.wo.as_f32().is_none()
13776                        || w.wq.rows() != self.num_heads * self.head_dim
13777                        || w.wk.rows() != self.num_kv_heads * self.head_dim
13778                        || w.wv.rows() != self.num_kv_heads * self.head_dim
13779                        || w.wq.cols() != self.hidden_size
13780                        || w.wk.cols() != self.hidden_size
13781                        || w.wv.cols() != self.hidden_size
13782                        || w.wo.rows() != self.hidden_size
13783                        || w.wo.cols() != self.num_heads * self.head_dim
13784                    {
13785                        return false;
13786                    }
13787                }
13788                _ => return false,
13789            }
13790            match &lw.ffn {
13791                FfnKind::Dense(d) => {
13792                    if d.act != Act::Silu
13793                        || !d.segs.is_empty()
13794                        || d.gate_proj.as_f32().is_none()
13795                        || d.up_proj.as_f32().is_none()
13796                        || d.down_proj.as_f32().is_none()
13797                        || d.gate_proj.rows() != self.intermediate_size
13798                        || d.gate_proj.cols() != self.hidden_size
13799                        || d.up_proj.rows() != self.intermediate_size
13800                        || d.up_proj.cols() != self.hidden_size
13801                        || d.down_proj.rows() != self.hidden_size
13802                        || d.down_proj.cols() != self.intermediate_size
13803                    {
13804                        return false;
13805                    }
13806                }
13807                FfnKind::Moe(m) => {
13808                    if m.resonance.is_none()
13809                        || m.top_k != 1
13810                        || m.router_sigmoid
13811                        || !m.norm_topk_prob
13812                        || m.expert_bias.is_some()
13813                        || m.routed_scaling != 1.0
13814                        || m.route_tau.is_some()
13815                        || m.shared.is_none()
13816                        || m.mask.is_some()
13817                        || m.per_expert_scale.is_some()
13818                        || m.router_input_norm
13819                        || m.experts.is_empty()
13820                        || m.experts.len() > 8
13821                    {
13822                        return false;
13823                    }
13824                    let r = m.resonance.as_ref().unwrap();
13825                    if r.mu.len() != m.experts.len() * self.hidden_size
13826                        || r.bias.len() != m.experts.len()
13827                        || r.u.len() != m.experts.len() * r.k * self.hidden_size
13828                        || r.k > 128
13829                    {
13830                        return false;
13831                    }
13832                    let Some((shared, gate)) = &m.shared else {
13833                        return false;
13834                    };
13835                    if gate.is_some() || shared.act != Act::Silu || !shared.segs.is_empty() {
13836                        return false;
13837                    }
13838                    if shared.gate_proj.as_f32().is_none()
13839                        || shared.up_proj.as_f32().is_none()
13840                        || shared.down_proj.as_f32().is_none()
13841                        || shared.gate_proj.rows() != self.intermediate_size
13842                        || shared.gate_proj.cols() != self.hidden_size
13843                        || shared.up_proj.rows() != self.intermediate_size
13844                        || shared.up_proj.cols() != self.hidden_size
13845                        || shared.down_proj.rows() != self.hidden_size
13846                        || shared.down_proj.cols() != self.intermediate_size
13847                    {
13848                        return false;
13849                    }
13850                    for e in &m.experts {
13851                        if e.act != Act::Silu
13852                            || !e.segs.is_empty()
13853                            || e.gate_proj.as_f32().is_none()
13854                            || e.up_proj.as_f32().is_none()
13855                            || e.down_proj.as_f32().is_none()
13856                            || e.gate_proj.rows() != self.intermediate_size
13857                            || e.gate_proj.cols() != self.hidden_size
13858                            || e.up_proj.rows() != self.intermediate_size
13859                            || e.up_proj.cols() != self.hidden_size
13860                            || e.down_proj.rows() != self.hidden_size
13861                            || e.down_proj.cols() != self.intermediate_size
13862                        {
13863                            return false;
13864                        }
13865                    }
13866                }
13867                FfnKind::DenseMoe(_) => return false,
13868            }
13869            if lw.input_norm.len() != self.hidden_size || lw.post_norm.len() != self.hidden_size {
13870                return false;
13871            }
13872        }
13873        if full_seen && self.anchor_core.is_some() {
13874            return false;
13875        }
13876        full_seen || self.num_layers > 0
13877    }
13878
13879    /// The resident Embryo graph is the owner of this pipeline's forward:
13880    /// the same gate `forward_layers_span` applies before handing a token
13881    /// to `forward_embryo_graph` (both graph phases on, the explicit
13882    /// opt-in, a wgpu device, no earlier refusal, an eligible stack).
13883    fn embryo_resident_wanted(&self) -> bool {
13884        crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
13885            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
13886            && matches!(
13887                std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
13888                Ok("1") | Ok("parallel")
13889            )
13890            && crate::gpu::enabled_here()
13891            && !self.graph_refused()
13892            && self.embryo_resident_eligible()
13893    }
13894
13895    /// Chunked prefill on the resident graph: `ids` from `start` in
13896    /// chunks of `EMBRYO_CHUNK_MAX`, one submit each, projections/FFN as
13897    /// chunk GEMMs and the recurrent layers walked in time on the device.
13898    /// Returns the last position's logits when the whole span ran there.
13899    /// `None` = refused before any device work (the per-position path
13900    /// takes the span).  `CMF_EMBRYO_CHUNK=0` keeps the per-position
13901    /// prefill (A/B and the parity reference).
13902    fn embryo_prefill_chunked(&mut self, ids: &[u32], start: usize) -> Option<Vec<f32>> {
13903        if ids.len() < 2
13904            || std::env::var("CMF_EMBRYO_CHUNK").as_deref() == Ok("0")
13905            || !self.embryo_resident_wanted()
13906        {
13907            return None;
13908        }
13909        let model = self.ensure_embryo_graph()?;
13910        let cmax = std::env::var("CMF_EMBRYO_CHUNK")
13911            .ok()
13912            .and_then(|v| v.parse::<usize>().ok())
13913            .filter(|&v| v >= 1)
13914            .unwrap_or(crate::gpu::EMBRYO_CHUNK_MAX)
13915            .min(crate::gpu::EMBRYO_CHUNK_MAX);
13916        let hs = self.hidden_size;
13917        let n = ids.len();
13918        let mut pos = start;
13919        let mut last = None;
13920        let mut rows = Vec::with_capacity(cmax * hs);
13921        while pos < n {
13922            let end = (pos + cmax).min(n);
13923            rows.clear();
13924            for &id in &ids[pos..end] {
13925                rows.extend_from_slice(&self.embed_single(id));
13926            }
13927            let mut lg = Vec::new();
13928            if !crate::gpu::forward_embryo_graph_chunk(
13929                &model,
13930                self.graph_kv_id,
13931                &rows,
13932                pos,
13933                end - pos,
13934                &mut lg,
13935            ) {
13936                if pos == start {
13937                    if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
13938                        eprintln!("embryo-dbg: chunk prefill refused at position {pos}");
13939                    }
13940                    return None;
13941                }
13942                // The device sequence advanced through the earlier chunks;
13943                // a host continuation would mix two owners of the state.
13944                // Fail this sequence alone, leaving no stale state or key.
13945                self.kv_cache.clear();
13946                self.clear_history();
13947                crate::gpu::graph_kv_reset(self.graph_kv_id);
13948                panic!(
13949                    "resident Embryo chunk prefill refused at position {pos}; refusing an in-flight host fallback"
13950                );
13951            }
13952            last = Some(lg);
13953            pos = end;
13954        }
13955        if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
13956            eprintln!(
13957                "embryo-dbg: chunk prefill {} ids from {start} in {} submits (chunk {cmax})",
13958                n - start,
13959                (n - start).div_ceil(cmax)
13960            );
13961        }
13962        last
13963    }
13964
13965    fn ensure_embryo_graph(&mut self) -> Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>> {
13966        if self.embryo_graph.is_none() && self.embryo_resident_eligible() {
13967            const UMAX: u32 = u32::MAX;
13968            const HEADER: usize = crate::gpu::EMBRYO_META_HEADER;
13969            const REC: usize = 64;
13970            struct Pack {
13971                data: Vec<f32>,
13972            }
13973            impl Pack {
13974                fn put(&mut self, x: &[f32]) -> u32 {
13975                    if x.is_empty() {
13976                        return u32::MAX;
13977                    }
13978                    let off = self.data.len();
13979                    self.data.extend_from_slice(x);
13980                    off as u32
13981                }
13982            }
13983            // The mixer family of the file: vmf_phase geometry fills the
13984            // phase header words, gated_delta_net fills words 24..29.  A
13985            // file has exactly one linear core, so at most one is live.
13986            let vmf = self.vmf_cfg;
13987            let gdn = self.gdn_cfg;
13988            let mut pack = Pack { data: Vec::new() };
13989            let mut meta = vec![0u32; HEADER];
13990            meta[0] = self.hidden_size as u32;
13991            meta[1] = self.intermediate_size as u32;
13992            meta[2] = self.vocab_size as u32;
13993            meta[3] = self.num_layers as u32;
13994            meta[4] = vmf.map(|c| c.num_heads).unwrap_or(0) as u32;
13995            meta[5] = vmf.map(|c| c.nphase).unwrap_or(0) as u32;
13996            meta[6] = vmf.map(|c| c.value_head_dim).unwrap_or(0) as u32;
13997            if let Some(g) = gdn {
13998                meta[24] = g.num_v_heads as u32;
13999                meta[25] = g.num_k_heads as u32;
14000                meta[26] = g.key_head_dim as u32;
14001                meta[27] = g.value_head_dim as u32;
14002                meta[28] = g.conv_kernel as u32;
14003                meta[29] = g.conv_dim() as u32;
14004            }
14005            meta[7] = self.num_heads as u32;
14006            meta[8] = self.num_kv_heads as u32;
14007            meta[9] = self.head_dim as u32;
14008            meta[10] = self.kv_cache.max_seq_len as u32;
14009            let clusters = self.head_clusters.as_ref().unwrap();
14010            let cluster_count = clusters.len() / self.hidden_size;
14011            if clusters.len() % self.hidden_size != 0
14012                || cluster_count == 0
14013                || cluster_count > 1024
14014                || self.vocab_size % cluster_count != 0
14015                || self.weights.lm_head.rows() < self.vocab_size
14016                || self.weights.final_norm.len() != self.hidden_size
14017            {
14018                return None;
14019            }
14020            meta[11] = cluster_count as u32;
14021            meta[12] = (self.vocab_size / cluster_count) as u32;
14022            meta[13] = vmf.map(|c| c.state_len()).unwrap_or(0) as u32;
14023            meta[16] = self.rotary_dim as u32;
14024            meta[17] = matches!(self.norm_style, NormStyle::Gemma) as u32;
14025            meta[19] = (self.rms_eps as f32).to_bits();
14026            let max_conv = self
14027                .weights
14028                .layers
14029                .iter()
14030                .filter_map(|lw| match &lw.attn {
14031                    AttnKind::Linear(w) => w.conv.as_ref().map(|v| v.len() / self.hidden_size),
14032                    _ => None,
14033                })
14034                .max()
14035                .unwrap_or(1);
14036            // One state slot per recurrent layer: the phase state plus its
14037            // hidden-wide conv ring, or the GDN record `[conv ring | S]`
14038            // (`GdnCfg::state_len`).  A file carries one mixer family, so
14039            // the stride is exactly that family's record and the device
14040            // state buffer equals the header's recurrent bytes.
14041            let phase_stride = vmf
14042                .map(|c| c.state_len() + max_conv.saturating_sub(1) * self.hidden_size)
14043                .unwrap_or(0);
14044            let gdn_stride = gdn.map(|g| g.state_len()).unwrap_or(0);
14045            let state_stride = phase_stride.max(gdn_stride);
14046            // Bounded genome: the KV plane of an anchor is its ring
14047            // `[kvh][W][hd]` K + V, and only anchors own one.  Legacy full
14048            // anchors keep the `max_seq` planes indexed by layer.
14049            let bounded = self.anchor_core.clone();
14050            let (anchor_window, anchor_sink) = bounded
14051                .as_ref()
14052                .map(|ac| (ac.window, ac.sink))
14053                .unwrap_or((0, 0));
14054            let kv_stride = if bounded.is_some() {
14055                2usize
14056                    .saturating_mul(self.num_kv_heads)
14057                    .saturating_mul(anchor_window)
14058                    .saturating_mul(self.head_dim)
14059            } else {
14060                2usize
14061                    .saturating_mul(self.num_kv_heads)
14062                    .saturating_mul(self.kv_cache.max_seq_len)
14063                    .saturating_mul(self.head_dim)
14064            };
14065            meta[14] = state_stride as u32;
14066            meta[15] = kv_stride as u32;
14067            meta[18] = anchor_window as u32;
14068            meta[20] = anchor_sink as u32;
14069            meta[21] = match &self.bounded_rope {
14070                Some(rope) => {
14071                    // [W][half] cos then [W][half] sin, one contiguous table.
14072                    let off = pack.put(&rope.cos);
14073                    let _ = pack.put(&rope.sin);
14074                    off
14075                }
14076                None => UMAX,
14077            };
14078            let mut full_seen = false;
14079            let mut bounded_seen = 0usize;
14080            // Recurrent state slots belong to mixer layers only (phase or
14081            // GDN): an anchor owns a ring, not a state stride, so the
14082            // device state buffer is exactly the header's recurrent bytes.
14083            let mut phase_seen = 0usize;
14084            let mut gdn_seen = 0usize;
14085            for (li, lw) in self.weights.layers.iter().enumerate() {
14086                let base = meta.len();
14087                meta.resize(base + REC, UMAX);
14088                meta[base] = match &lw.attn {
14089                    AttnKind::Linear(w) if w.phase_delta => 1,
14090                    AttnKind::Linear(_) => 0,
14091                    AttnKind::Full { .. } => 2,
14092                    AttnKind::Bounded(_) => 3,
14093                    AttnKind::LinearGdn(_) => 4,
14094                    _ => UMAX,
14095                };
14096                meta[base + 1] = pack.put(&lw.input_norm);
14097                meta[base + 2] = pack.put(&lw.post_norm);
14098                meta[base + 25] = match &lw.attn {
14099                    AttnKind::Linear(_) => {
14100                        let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14101                        phase_seen += 1;
14102                        off
14103                    }
14104                    AttnKind::LinearGdn(_) => {
14105                        let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14106                        gdn_seen += 1;
14107                        off
14108                    }
14109                    _ => UMAX,
14110                };
14111                match &lw.attn {
14112                    AttnKind::LinearGdn(w) => {
14113                        // Layer record words 56..63 + 29, as the resident
14114                        // kernels read them (gpu_wgpu.rs `embryo_core_gdn_*`).
14115                        meta[base + 56] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_qkv));
14116                        meta[base + 57] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_z));
14117                        meta[base + 58] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_a));
14118                        meta[base + 59] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_b));
14119                        meta[base + 60] = pack.put(&w.conv1d);
14120                        meta[base + 61] = pack.put(&w.a_log);
14121                        meta[base + 62] = pack.put(&w.dt_bias);
14122                        meta[base + 63] = pack.put(&w.norm);
14123                        meta[base + 29] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14124                        meta[base + 24] = 0;
14125                    }
14126                    AttnKind::Linear(w) => {
14127                        meta[base + 3] = pack.put(&Self::embryo_qtensor_f32(&w.thq));
14128                        meta[base + 4] = pack.put(&Self::embryo_qtensor_f32(&w.thk));
14129                        meta[base + 5] = pack.put(&Self::embryo_qtensor_f32(&w.v_proj));
14130                        meta[base + 6] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14131                        let decay: Vec<f32> = w.decay.iter().map(|&x| x as f32).collect();
14132                        meta[base + 7] = pack.put(&decay);
14133                        if let Some((kg, kb)) = &w.k_gate {
14134                            meta[base + 8] = pack.put(&Self::embryo_qtensor_f32(kg));
14135                            meta[base + 9] = pack.put(kb);
14136                        }
14137                        if let Some(conv) = &w.conv {
14138                            meta[base + 10] = pack.put(conv);
14139                            meta[base + 24] = (conv.len() / self.hidden_size) as u32;
14140                        } else {
14141                            meta[base + 24] = 0;
14142                        }
14143                    }
14144                    AttnKind::Full { wq, wk, wv, wo, .. } => {
14145                        full_seen = true;
14146                        meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(wq));
14147                        meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(wk));
14148                        meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(wv));
14149                        meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(wo));
14150                        meta[base + 26] = (li * kv_stride) as u32;
14151                    }
14152                    AttnKind::Bounded(w) => {
14153                        meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(&w.wq));
14154                        meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(&w.wk));
14155                        meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(&w.wv));
14156                        meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(&w.wo));
14157                        // Ring slot of this anchor (anchors only, packed).
14158                        meta[base + 26] = (bounded_seen * kv_stride) as u32;
14159                        meta[base + 27] = pack.put(&w.sink_k);
14160                        meta[base + 28] = pack.put(&w.sink_v);
14161                        bounded_seen += 1;
14162                    }
14163                    _ => return None,
14164                }
14165                match &lw.ffn {
14166                    FfnKind::Dense(d) => {
14167                        meta[base + 15] = 0;
14168                        meta[base + 16] = 0;
14169                        meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&d.gate_proj));
14170                        meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&d.up_proj));
14171                        meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&d.down_proj));
14172                    }
14173                    FfnKind::Moe(m) => {
14174                        let r = m.resonance.as_ref().unwrap();
14175                        let (shared, _) = m.shared.as_ref().unwrap();
14176                        meta[base + 15] = 1;
14177                        meta[base + 16] = m.experts.len() as u32;
14178                        meta[base + 17] = pack.put(&r.mu);
14179                        meta[base + 18] = pack.put(&r.u);
14180                        meta[base + 19] = pack.put(&r.bias);
14181                        meta[base + 20] = r.k as u32;
14182                        // Word 30: the growth shell as the runtime applies
14183                        // it now (`+inf` on trunk rows, the stored finite
14184                        // shell on grown rows, all `+inf` under
14185                        // `CMF_GROWTH_SHELL=off`), followed by one `−∞`
14186                        // sentinel at index E the kernel writes as the
14187                        // score of an expert outside its shell — WGSL has
14188                        // no infinity literal, so the value travels as
14189                        // data (`embryo_core_route_finalize`).
14190                        let mut shell = r.effective_shell(m.experts.len());
14191                        shell.push(f32::NEG_INFINITY);
14192                        meta[base + 30] = pack.put(&shell);
14193                        meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&shared.gate_proj));
14194                        meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&shared.up_proj));
14195                        meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&shared.down_proj));
14196                        for (e, ex) in m.experts.iter().enumerate() {
14197                            meta[base + 32 + e * 3] =
14198                                pack.put(&Self::embryo_qtensor_f32(&ex.gate_proj));
14199                            meta[base + 33 + e * 3] =
14200                                pack.put(&Self::embryo_qtensor_f32(&ex.up_proj));
14201                            meta[base + 34 + e * 3] =
14202                                pack.put(&Self::embryo_qtensor_f32(&ex.down_proj));
14203                        }
14204                    }
14205                    FfnKind::DenseMoe(_) => return None,
14206                }
14207            }
14208            if !full_seen && self.num_layers == 0 {
14209                return None;
14210            }
14211            let id = {
14212                static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
14213                NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
14214            };
14215            let model = crate::gpu::EmbryoGraphModel {
14216                id,
14217                hidden: self.hidden_size,
14218                intermediate: self.intermediate_size,
14219                vocab: self.vocab_size,
14220                layers: self.num_layers,
14221                phase_heads: vmf.map(|c| c.num_heads).unwrap_or(0),
14222                nphase: vmf.map(|c| c.nphase).unwrap_or(0),
14223                phase_dv: vmf.map(|c| c.value_head_dim).unwrap_or(0),
14224                anchor_q_heads: self.num_heads,
14225                anchor_kv_heads: self.num_kv_heads,
14226                anchor_head_dim: self.head_dim,
14227                rotary_dim: self.rotary_dim,
14228                max_seq: self.kv_cache.max_seq_len,
14229                cluster_count,
14230                cluster_size: self.vocab_size / cluster_count,
14231                phase_state_len: vmf.map(|c| c.state_len()).unwrap_or(0),
14232                state_stride,
14233                kv_stride,
14234                norm_gemma: matches!(self.norm_style, NormStyle::Gemma),
14235                phase_mass: vmf.map(|c| c.phase_mass).unwrap_or(0.0),
14236                weights: pack.data,
14237                meta,
14238                lm_head: Self::embryo_qtensor_f32(&self.weights.lm_head),
14239                clusters: clusters.as_ref().clone(),
14240                final_norm: self.weights.final_norm.clone(),
14241                inv_freq: self.inv_freq.as_ref().clone(),
14242                bounded: bounded.is_some(),
14243                kv_layers: if bounded.is_some() {
14244                    bounded_seen
14245                } else {
14246                    self.num_layers
14247                },
14248                state_layers: phase_seen + gdn_seen,
14249                anchor_window,
14250                anchor_sink,
14251                phase_layers: phase_seen,
14252                gdn_layers: gdn_seen,
14253                gdn_heads: gdn.map(|g| g.num_v_heads).unwrap_or(0),
14254                gdn_k_heads: gdn.map(|g| g.num_k_heads).unwrap_or(0),
14255                gdn_dk: gdn.map(|g| g.key_head_dim).unwrap_or(0),
14256                gdn_dv: gdn.map(|g| g.value_head_dim).unwrap_or(0),
14257                gdn_kk: gdn.map(|g| g.conv_kernel).unwrap_or(0),
14258            };
14259            self.embryo_graph = Some(std::sync::Arc::new(model));
14260        }
14261        self.embryo_graph.clone()
14262    }
14263
14264    fn forward_layers_span(
14265        &mut self,
14266        hidden: &[f32],
14267        position: usize,
14268        task_mask: Option<&TaskMask>,
14269        from: usize,
14270        upto: Option<usize>,
14271    ) -> Vec<f32> {
14272        debug_assert!(
14273            from == 0
14274                || (self.dsv4.is_none()
14275                    && self.dsv41.is_none()
14276                    && self.qwen4_exp.is_none()
14277                    && self.g3n.is_none())
14278        );
14279        // Every plain forward — the whole-token Metal graph (`q1_graph_gpu`
14280        // wraps the GDN owners zero-copy and reallocates them on a size
14281        // change) and the CPU layer loop (reads/swaps `linear_state`) —
14282        // must see the previous speculative commit's asynchronous replay
14283        // complete. One mutex probe when nothing is pending.
14284        #[cfg(target_os = "macos")]
14285        if !crate::gpu_metal::wait_replay() {
14286            self.fail_metal_graph("the pending async replay failed before a plain forward");
14287            return vec![0.0; self.hidden_size];
14288        }
14289        if let Some(b) = &mut self.qwen4_exp {
14290            let _ = (task_mask, upto);
14291            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14292            let mut logits = Vec::new();
14293            crate::qwen4_exp::forward_token(
14294                &b.0,
14295                &b.1,
14296                &b.2,
14297                &mut b.3,
14298                token_id,
14299                position,
14300                &self.inv_freq,
14301                self.pool.as_deref(),
14302                &mut logits,
14303                true,
14304            );
14305            self.graph_logits = Some(logits);
14306            return vec![0.0; self.hidden_size];
14307        }
14308        // DeepSeek-V4 runs its own stack: the state is hc_mult copies, and
14309        // the forward returns LOGITS, not a hidden — the head is inside it
14310        // (the final fold sits between the last layer and the norm). The
14311        // token id rides in `hidden[0]`, written by embed_single, because
14312        // the hash layers route by id rather than by content.
14313        if let Some(b) = &mut self.dsv4 {
14314            let _ = (task_mask, upto);
14315            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14316            let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
14317            st.pos = position;
14318            let mut logits = Vec::new();
14319            crate::dsv4::forward_token(
14320                g,
14321                layers,
14322                &cfg,
14323                st,
14324                token_id,
14325                &self.inv_freq,
14326                self.pool.as_deref(),
14327                &mut logits,
14328            );
14329            self.graph_logits = Some(logits);
14330            self.dspark_probe(position, token_id);
14331            // The caller expects a hidden; the logits went out of band, as
14332            // with the fused lm_head path.
14333            return vec![0.0; self.hidden_size];
14334        }
14335        // DeepSeek-V4.1 owns its complete stack and emits logits out of band.
14336        if let Some(b) = &mut self.dsv41 {
14337            let _ = (task_mask, upto);
14338            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14339            let mut logits = Vec::new();
14340            crate::dsv41::forward_token(
14341                &b.0,
14342                &b.1,
14343                &b.2,
14344                &mut b.3,
14345                token_id,
14346                position,
14347                self.pool.as_deref(),
14348                &mut logits,
14349            );
14350            self.graph_logits = Some(logits);
14351            return vec![0.0; self.hidden_size];
14352        }
14353        // Gemma-3n runs its own stack (4 AltUp replicas don't fit this
14354        // loop); `hidden` is the extended embedding from embed_single.
14355        if let Some(b) = &self.g3n {
14356            let _ = (task_mask, upto);
14357            return crate::g3n::g3n_forward(
14358                &b.0,
14359                &b.1,
14360                hidden,
14361                position,
14362                &mut self.kv_cache.layers,
14363                self.num_heads,
14364                self.num_kv_heads,
14365                self.head_dim,
14366                self.pool.as_deref(),
14367            );
14368        }
14369        // Cortiq Embryo owns a separate resident graph: phase recurrent
14370        // state, resonance routing, the GQA anchor KV and hierarchical head
14371        // all execute in one Vulkan submit. It is limited to a complete
14372        // unmasked stack; spans and task masks retain the exact host path.
14373        // `CMF_EMBRYO_DBG=1` names the gate that keeps a token off the
14374        // resident graph — every refusal below is otherwise silent.
14375        if from == 0
14376            && upto.is_none()
14377            && task_mask.is_none()
14378            && self.anchor_core.is_some()
14379            && std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1")
14380        {
14381            static ONCE: std::sync::Once = std::sync::Once::new();
14382            ONCE.call_once(|| {
14383                eprintln!(
14384                    "embryo-dbg: graph_on decode={} prefill={} resident_env={:?} enabled_here={} \
14385                     unsupported={} eligible={}",
14386                    crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode),
14387                    crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill),
14388                    std::env::var("CMF_EMBRYO_RESIDENT").ok(),
14389                    crate::gpu::enabled_here(),
14390                    self.graph_refused(),
14391                    self.embryo_resident_eligible(),
14392                );
14393            });
14394        }
14395        if from == 0
14396            && upto.is_none()
14397            && task_mask.is_none()
14398            // Embryo's recurrent/KV state has no host import path.  Do not
14399            // seed it for a prefill-only graph and then silently decode from
14400            // an empty CPU cache; both phases must select the resident owner.
14401            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14402            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14403            // This whole-token Embryo path remains explicitly opt-in.
14404            // `CMF_GPU_WGPU_GRAPH=1` still enables the mature generic graph,
14405            // but must not silently select this model-specific resident path.
14406            && matches!(
14407                std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14408                Ok("1") | Ok("parallel")
14409            )
14410            && crate::gpu::enabled_here()
14411            && !self.graph_refused()
14412            // Sequence owner: past position zero the device continues only
14413            // a sequence it holds. A host-owned sequence (the graph refused
14414            // at its start, or its prefix was prefilled on the host) keeps
14415            // the host path to its end — never a device attempt at p > 0
14416            // over an empty device image.
14417            && (position == 0 || self.device_sequence_position().is_some())
14418            && self.embryo_resident_eligible()
14419            && let Some(model) = self.ensure_embryo_graph()
14420        {
14421            let mut lg = Vec::new();
14422            if crate::gpu::forward_embryo_graph(&model, self.graph_kv_id, hidden, position, &mut lg)
14423            {
14424                self.graph_logits = Some(lg);
14425                return vec![0.0; self.hidden_size];
14426            }
14427            if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14428                eprintln!("embryo-dbg: forward_embryo_graph refused at position {position}");
14429            }
14430            // The refusal is THIS pipeline's: falling through once is safe
14431            // at position zero (the host owns the sequence from here),
14432            // while a refusal at a later position of a device-owned
14433            // sequence would mix a host KV/state path with a partial
14434            // device sequence.
14435            self.mark_graph_refused();
14436            if position != 0 {
14437                // Fail this sequence alone and leave nothing stale behind:
14438                // no reuse key, no host or device state for the next
14439                // request on this slot to "extend".
14440                self.kv_cache.clear();
14441                self.clear_history();
14442                crate::gpu::graph_kv_reset(self.graph_kv_id);
14443                panic!(
14444                    "resident Embryo graph refused at position {position}; refusing an in-flight host fallback"
14445                );
14446            }
14447        }
14448        let mut h = hidden.to_vec();
14449        // MiMo-V2 expert placement: decided before the graph or the per-op
14450        // arena can claim the budget the expert bank needs.
14451        self.mimo_moe_prepare();
14452        let _mimo_q8 = self.mimo_moe.is_on()
14453            .then(crate::qtensor::enter_full_gpu_q8_scope);
14454        // Split borrows: copy scalars / clone handles so the per-layer
14455        // cfg does not hold `&self` while the KV cache is `&mut`.
14456        let (nh, _nkv, _hd, hs, _rd, eps) = (
14457            self.num_heads,
14458            self.num_kv_heads,
14459            self.head_dim,
14460            self.hidden_size,
14461            self.rotary_dim,
14462            self.rms_eps,
14463        );
14464        let pool = self.pool.clone();
14465        // Opt-in wgpu token-graph attention (discrete Vulkan/DX12): the whole
14466        // attention sub-block runs resident in one submit. Off by default.
14467        // Whole-token wgpu graph: eligibility + arbitration.
14468        //  - explicit CMF_GPU_WGPU_GRAPH forces it on/off;
14469        //  - discrete adapters (4090: decode 76 -> 137 tok/s) and GDN
14470        //    hybrids (recurrent state device-resident, no CPU twin to
14471        //    race) TRUST it;
14472        //  - integrated/mobile adapters RACE it against the normal path
14473        //    at generation granularity (gpu::graph_race_*) — tiled
14474        //    mobile GPUs can turn the ~300-dispatch graph into seconds
14475        //    per token, while a fast phone GPU keeps its win.
14476        let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
14477        let graph_on = match graph_env.as_deref() {
14478            Some("0") => false,
14479            Some("prefill") => false, // decode keeps the per-op path
14480            Some(_) => true,
14481            // Unset: same discrete-only default as every other graph
14482            // site. "Is the GPU on" used to stand in here — which made
14483            // the 0.2 tok/s whole-token graph race-eligible on mobile
14484            // adapters and cost 12-14× on first tokens (cmfmobile
14485            // TUNING.md); integrated GPUs keep the per-op probe path.
14486            None => crate::gpu::wgpu_graph_default(),
14487        };
14488        let graph_trusted =
14489            graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
14490        let race_eligible = graph_on
14491            && upto.is_none()
14492            && task_mask.is_none()
14493            && from == 0
14494            && !self.graph_refused();
14495        let mut tail_start = 0usize;
14496        if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
14497            let t_graph = std::time::Instant::now();
14498            let mut lg = Vec::new();
14499            let mut gl = 0usize;
14500            let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
14501            let declined = built.is_none();
14502            let built = match built {
14503                Some(Ok(hh)) => Some(hh),
14504                Some(Err(())) => {
14505                    // O(1) state was admitted before the device failure; the
14506                    // CPU mirrors are stale by construction.  Clear the whole
14507                    // sequence and stop rather than walking that stale state.
14508                    self.clear_sequence_state();
14509                    self.graph_failed
14510                        .store(true, std::sync::atomic::Ordering::Relaxed);
14511                    self.cancel
14512                        .store(true, std::sync::atomic::Ordering::Relaxed);
14513                    tracing::error!("token graph failed after admission; sequence state cleared");
14514                    return vec![0.0; self.hidden_size];
14515                }
14516                None => None,
14517            };
14518            // Past the transient guards (o1 still collecting, a softcap)
14519            // a refusal is about the weights and will never change —
14520            // remember it instead of walking every layer again next
14521            // token.
14522            if declined && !self.o1_active() && self.attn_softcap == 0.0 {
14523                self.mark_graph_refused();
14524            }
14525            graph_note(built.is_some(), gl, self.num_layers);
14526            if let Some(hh) = built {
14527                let dur = t_graph.elapsed();
14528                if std::env::var("CMF_GRAPH_PROF").is_ok() {
14529                    eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
14530                }
14531                if gl > 0 && gl < self.num_layers {
14532                    // Device prefix: the graph ran layers 0..gl and handed
14533                    // back the boundary hidden — the loop below owns the
14534                    // tail. The prefix layers' KV/state advanced on the
14535                    // device; the tail's advances on the host below. One
14536                    // boundary crossing per token.
14537                    h = hh;
14538                    tail_start = gl;
14539                } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
14540                    if !graph_trusted {
14541                        crate::gpu::graph_race_record(true, dur);
14542                    }
14543                    if !lg.is_empty() {
14544                        // Graph produced logits (final-norm + lm_head folded in) —
14545                        // pad/cap to vocab and hand them to the sampler directly.
14546                        lg.resize(self.vocab_size, 0.0);
14547                        if let Some(c) = self.final_softcap {
14548                            for l in lg.iter_mut() {
14549                                *l = c * (*l / c).tanh();
14550                            }
14551                        }
14552                        self.graph_logits = Some(lg);
14553                    }
14554                    return hh;
14555                }
14556                // Hopeless first graph token: discard it and fall through
14557                // to the normal path. Safe exactly here — the prompt KV is
14558                // still CPU-owned (chunked prefill), so recomputing this
14559                // position is exact; the mirror's extra row is never read
14560                // (the race just settled on the normal path).
14561            }
14562        }
14563        // KIMI-LINEAR HAS NO SPLIT BUG. The 2.6× reported from the
14564        // model rotation (12.2 tok/s on one card against 4.6 on two)
14565        // was a single measurement of a model whose arm arbitration is
14566        // borderline, and it did not survive repetition. Three runs an
14567        // arm, same binary, back to back:
14568        //   probe on : 1 GPU 9.5 / 5.7 / 5.9   2 GPU 7.8 / 13.0 / 13.3
14569        //   pinned   : 1 GPU 5.6 / 5.3 / 5.2   2 GPU 3.5 / 4.2 / 3.4
14570        // With the arms pinned the split costs about 1.45×, which is
14571        // what a layer split costs. With the probe free, TWO CARDS RUN
14572        // FASTER — because for this model the CPU arm wins some op
14573        // classes and the probe finds that.
14574        //
14575        // Two things do stand, and both are measured. The token graph
14576        // builds NOTHING here (`covered 0 of 14 layers [0..14)`), so
14577        // every layer walks per-op on either arm — that is where the
14578        // headroom is, not in the split. And this model's benchmark is
14579        // unusable without `CMF_GPU_PROBE=0`: the arbitration alone
14580        // moves it by more than 2×.
14581        //
14582        // Span runs (network split): the graph covers exactly [from..=upto]
14583        // — one submit per SEGMENT per token. No race: its state is global
14584        // and calibrated on full stacks, so spans take the graph only where
14585        // it is trusted by default (discrete adapters / CMF_GPU_WGPU_GRAPH).
14586        let span = from > 0 || upto.is_some();
14587        if span && graph_on && task_mask.is_none() && graph_trusted {
14588            let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
14589            let mut lg = Vec::new();
14590            let mut gl = 0usize;
14591            let span_res =
14592                self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
14593            let span_res = match span_res {
14594                Some(Ok(hh)) => Some(hh),
14595                Some(Err(())) => {
14596                    self.clear_sequence_state();
14597                    self.graph_failed
14598                        .store(true, std::sync::atomic::Ordering::Relaxed);
14599                    self.cancel
14600                        .store(true, std::sync::atomic::Ordering::Relaxed);
14601                    tracing::error!(
14602                        "span token graph failed after admission; sequence state cleared"
14603                    );
14604                    return vec![0.0; self.hidden_size];
14605                }
14606                None => None,
14607            };
14608            graph_note(span_res.is_some(), gl, upto_excl - from);
14609            if std::env::var("CMF_GPU_DEBUG").is_ok() {
14610                // How much of the span the graph actually covered. A
14611                // prefix of nothing means every layer walks per-op and
14612                // the split's extra cost is elsewhere.
14613                static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
14614                if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
14615                    eprintln!(
14616                        "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
14617                        upto_excl - from,
14618                        span_res.is_some()
14619                    );
14620                }
14621            }
14622            if let Some(hh) = span_res {
14623                if gl == upto_excl - from {
14624                    if !lg.is_empty() {
14625                        lg.resize(self.vocab_size, 0.0);
14626                        if let Some(c) = self.final_softcap {
14627                            for l in lg.iter_mut() {
14628                                *l = c * (*l / c).tanh();
14629                            }
14630                        }
14631                        self.graph_logits = Some(lg);
14632                    }
14633                    crate::gpu::set_layer(-1);
14634                    return hh;
14635                }
14636                // Partial device prefix of the span: CPU owns the tail.
14637                h = hh;
14638                tail_start = from + gl;
14639            }
14640        }
14641        // Layers the host is about to run whose device mirror moved ahead
14642        // of the host cache (a device prefix that shrank since the prompt,
14643        // a batched-prefill prefix longer than this token's): bring their
14644        // rows over first. One comparison per layer when nothing lags.
14645        let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
14646
14647        // A partial graph is an explicit GPU-prefix / CPU-tail split. Keep
14648        // the tail PURE host-side: letting its QTensor hooks re-enter the
14649        // residency arena streams every omitted layer through Vulkan and the
14650        // driver's freed-allocation cache can grow to the full model size
14651        // (25.4 GiB observed with a 14 GiB budget on Granite 30B Q8_2F).
14652        // With a MiMo expert bank the tail is not a whole-layer host
14653        // stream: its experts run from the bank (never the arena) and its
14654        // projections stay per-op on the device, which the bank's placement
14655        // left room for.
14656        let host_tail = tail_start > from;
14657        let _host_tail = (host_tail && !self.mimo_moe.is_on()).then(crate::gpu::enter_cpu_scope);
14658        let automatic_gpu_prefix = self.automatic_gpu_prefix();
14659
14660        let _prof_layers = crate::cpuprof::time(crate::cpuprof::Slot::Layers);
14661        #[cfg(target_os = "macos")]
14662        let mut gpu_skip_until = 0usize;
14663        for li in tail_start.max(from)..self.num_layers {
14664            let _capacity_tail = automatic_gpu_prefix
14665                .filter(|&prefix| li >= prefix && !self.mimo_moe.is_dynamic(li, host_tail))
14666                .map(|_| crate::gpu::enter_cpu_scope());
14667            crate::gpu::set_layer(li as i64); // layer-split GPU/CPU (CMF_GPU_LAYERS)
14668            if let Some(u) = upto {
14669                if li > u {
14670                    break;
14671                }
14672            }
14673            if let Some(mask) = task_mask {
14674                if !mask.layer_alive(li) {
14675                    continue; // dead layer: residual pass-through
14676                }
14677            }
14678            // Whole-block q1 token graph: a run of consecutive q1
14679            // layers — GDN and full attention — executes with one sync
14680            // per CPU attend instead of per op (macOS/Metal).
14681            #[cfg(target_os = "macos")]
14682            {
14683                if li < gpu_skip_until {
14684                    continue;
14685                }
14686                if task_mask.is_none() {
14687                    let end = self.q1_graph_gpu(li, upto, position, &mut h);
14688                    if self
14689                        .graph_failed
14690                        .load(std::sync::atomic::Ordering::Relaxed)
14691                    {
14692                        // The graph may have mutated device state before a
14693                        // command-buffer error. Never continue with a CPU
14694                        // tail or read a stale host mirror after admission.
14695                        return vec![0.0; self.hidden_size];
14696                    }
14697                    if end > li {
14698                        gpu_skip_until = end;
14699                        // Looped Transformer: the graph stopped at a loop
14700                        // boundary — apply final norm before the next iteration.
14701                        if self.is_loop_end(end - 1) && end < self.num_layers {
14702                            h = inference::rms_norm(
14703                                &h,
14704                                &self.weights.final_norm,
14705                                self.rms_eps,
14706                                self.norm_style,
14707                            );
14708                        }
14709                        continue;
14710                    }
14711                }
14712            }
14713
14714            if task_mask.is_none() {
14715                match self.mimo_graph_layer_rows(li, &mut h, &[position]) {
14716                    crate::gpu::BatchGraphOutcome::Completed => continue,
14717                    crate::gpu::BatchGraphOutcome::Failed => return vec![0.0; self.hidden_size],
14718                    crate::gpu::BatchGraphOutcome::Declined => {},
14719                }
14720            }
14721            #[cfg(feature = "gpu")]
14722            self.pull_lagging_host_kv(li, li + 1, position);
14723            let lw = &self.weights.layers[self.phys_layer(li)];
14724            if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
14725                if tp.parse::<usize>().ok() == Some(position) {
14726                    let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
14727                    eprintln!(
14728                        "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
14729                        h[0], h[1]
14730                    );
14731                }
14732            }
14733            // Norm into the pipeline scratch — the returning rms_norm
14734            // allocated twice per layer per token (roadmap §3 P0).
14735            let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
14736            inference::rms_norm_into(
14737                &h,
14738                &lw.input_norm,
14739                self.rms_eps,
14740                self.norm_style,
14741                &mut self.ws.n1,
14742            );
14743            drop(prof);
14744
14745            let attn_out = match &lw.attn {
14746                AttnKind::Mla(w) => {
14747                    let inv_freq_l = self.layer_inv_freq(li);
14748                    let rs = self.layer_rope_scale(li);
14749                    let eps = self.rms_eps;
14750                    let pool = self.pool.clone();
14751                    mla_attention(
14752                        w,
14753                        &self.ws.n1,
14754                        &mut self.kv_cache.layers[li],
14755                        position,
14756                        &inv_freq_l,
14757                        rs,
14758                        eps,
14759                        pool.as_deref(),
14760                    )
14761                }
14762                AttnKind::Linear(w) => {
14763                    let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
14764                    vmf_phase_forward(
14765                        &self.ws.n1,
14766                        w,
14767                        &cfg,
14768                        &mut self.kv_cache.layers[li].linear_state,
14769                        self.pool.as_deref(),
14770                    )
14771                }
14772                AttnKind::Kda(w) => {
14773                    let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
14774                    crate::linear_core::kda_forward(
14775                        &self.ws.n1,
14776                        w,
14777                        &cfg,
14778                        &mut self.kv_cache.layers[li].linear_state,
14779                        self.pool.as_deref(),
14780                    )
14781                }
14782                AttnKind::LinearGdn(w) => {
14783                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
14784                    gdn_forward(
14785                        &self.ws.n1,
14786                        w,
14787                        &cfg,
14788                        &mut self.kv_cache.layers[li].linear_state,
14789                        self.pool.as_deref(),
14790                    )
14791                }
14792                AttnKind::ShortConv(w) => {
14793                    let cfg = self
14794                        .short_conv_cfg
14795                        .expect("short-conv layer without short_conv_cfg");
14796                    short_conv_forward(
14797                        &self.ws.n1,
14798                        w,
14799                        &cfg,
14800                        &mut self.kv_cache.layers[li].linear_state,
14801                        self.pool.as_deref(),
14802                    )
14803                }
14804                AttnKind::Bounded(w) => {
14805                    // Natively bounded anchor: insert into the ring, attend
14806                    // over sinks ∪ window. No position, nothing appended.
14807                    let rope = self
14808                        .bounded_rope
14809                        .clone()
14810                        .expect("bounded layer without an installed rotation table");
14811                    let cfg = crate::bounded::BoundedAttnCfg {
14812                        num_heads: self.num_heads,
14813                        num_kv_heads: self.num_kv_heads,
14814                        head_dim: self.head_dim,
14815                        hidden_size: hs,
14816                        scale: self.attn_scale,
14817                        rope: &rope,
14818                        pool: pool.as_deref(),
14819                    };
14820                    crate::bounded::bounded_attention(
14821                        &self.ws.n1,
14822                        w,
14823                        &mut self.kv_cache.layers[li],
14824                        &cfg,
14825                    )
14826                }
14827                AttnKind::Full {
14828                    wq,
14829                    wk,
14830                    wv,
14831                    wo,
14832                    q_norm,
14833                    k_norm,
14834                    output_gate,
14835                    softplus_gate,
14836                    bias,
14837                } if self.kv_cache.layers[li].o1_sealed() => {
14838                    // O(1) override: decode on the sealed Nyström state
14839                    // instead of the growing KV cache.
14840                    let inv_freq_l = self.layer_inv_freq(li);
14841                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14842                    let cfg = QwenAttnCfg {
14843                        num_heads: self.layer_num_heads(li),
14844                        num_kv_heads: nkv_l,
14845                        head_dim: hd_l,
14846                        hidden_size: hs,
14847                        position,
14848                        inv_freq: &inv_freq_l,
14849                        rotary_dim: rd_l,
14850                        scale: self.attn_scale,
14851                        softcap: self.attn_softcap,
14852                        window: None,
14853                        v_norm: self.attn_v_norm,
14854                        qk_norm_after_rope: self.qk_norm_after_rope,
14855                        q_norm: q_norm.as_deref(),
14856                        k_norm: k_norm.as_deref(),
14857                        output_gate: *output_gate,
14858                        softplus_gate: softplus_gate
14859                            .as_ref()
14860                            .map(|(gate, per_head)| (gate, *per_head)),
14861                        rope_scale: self.layer_rope_scale(li),
14862                        bias: bias
14863                            .as_ref()
14864                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
14865                        rms_eps: eps,
14866                        norm_style: self.norm_style,
14867                        pool: pool.as_deref(),
14868                        v_head_dim: self.layer_v_dim(li),
14869                    };
14870                    attention::qwen_attention_nystrom(
14871                        &self.ws.n1,
14872                        wq,
14873                        wk,
14874                        wv,
14875                        wo,
14876                        &mut self.kv_cache.layers[li],
14877                        &cfg,
14878                    )
14879                }
14880                AttnKind::Full {
14881                    wq,
14882                    wk,
14883                    wv,
14884                    wo,
14885                    q_norm,
14886                    k_norm,
14887                    output_gate,
14888                    softplus_gate,
14889                    bias,
14890                } => 'attn: {
14891                    // wgpu token-graph attention (opt-in): whole sub-block in
14892                    // one submit, device K/V mirror. q1 only, no gate/bias/mask.
14893                    // Its kernel has no window, sink or narrow-V slot and
14894                    // one mirror geometry: such models stay on the CPU attend.
14895                    let dropin_reason =
14896                        graph_on.then(|| self.graph_attn_decline_reason()).flatten();
14897                    if let Some(reason) = dropin_reason {
14898                        self.note_graph_decline("wgpu attn dropin", reason);
14899                    }
14900                    if graph_on
14901                        && dropin_reason.is_none()
14902                        && !*output_gate
14903                        && softplus_gate.is_none()
14904                        && self.attention_heads_per_layer.is_none()
14905                        && bias.is_none()
14906                        && task_mask.is_none()
14907                    {
14908                        let inv_freq_l = self.layer_inv_freq(li);
14909                        let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14910                        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
14911                        if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
14912                            wq.mapped_q1(),
14913                            wk.mapped_q1(),
14914                            wv.mapped_q1(),
14915                            wo.mapped_q1(),
14916                        ) {
14917                            let gm = gm.clone();
14918                            let mut out = vec![0f32; hs];
14919                            let cache = &self.kv_cache.layers[li];
14920                            if crate::gpu::attn_dropin(
14921                                &gm,
14922                                self.graph_kv_id,
14923                                li,
14924                                &self.ws.n1,
14925                                qi,
14926                                ki,
14927                                vi,
14928                                oi,
14929                                q_norm.as_deref(),
14930                                k_norm.as_deref(),
14931                                self.qk_norm_after_rope,
14932                                &inv_freq_l,
14933                                nh,
14934                                nkv_l,
14935                                hd_l,
14936                                rd_l,
14937                                hs,
14938                                position,
14939                                self.kv_cache.max_seq_len,
14940                                gemma,
14941                                eps as f32,
14942                                cache.k_heads(),
14943                                cache.v_heads(),
14944                                &mut out,
14945                            ) {
14946                                break 'attn out;
14947                            }
14948                        }
14949                    }
14950                    let masked = task_mask
14951                        .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
14952                        .unwrap_or(false);
14953                    let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
14954                    // The masked kernel knows one pipeline-wide geometry and
14955                    // RoPE table, no window and no sink.
14956                    let plain = self.layer_attn_plain(li);
14957                    match (masked, f32_view) {
14958                        // Historical masked path (f32 slices; the loader
14959                        // keeps masked models in f32).
14960                        (true, (Some(q), Some(k), Some(v), Some(o))) if plain => {
14961                            let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
14962                            attention::multi_head_attention(
14963                                &self.ws.n1,
14964                                q,
14965                                k,
14966                                v,
14967                                o,
14968                                &mut self.kv_cache.layers[li],
14969                                self.num_heads,
14970                                self.num_kv_heads,
14971                                self.head_dim,
14972                                self.hidden_size,
14973                                position,
14974                                &active_heads,
14975                                &self.inv_freq,
14976                            )
14977                        }
14978                        (masked, _) => {
14979                            if masked {
14980                                tracing::warn!(
14981                                    "layer {li}: head mask on quantized weights or on a \
14982                                     window/sink/per-layer-geometry layer not supported \
14983                                     yet — executing dense"
14984                                );
14985                            }
14986                            let inv_freq_l = self.layer_inv_freq(li);
14987                            let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
14988                            let cfg = QwenAttnCfg {
14989                                num_heads: self.layer_num_heads(li),
14990                                num_kv_heads: nkv_l,
14991                                head_dim: hd_l,
14992                                hidden_size: hs,
14993                                position,
14994                                inv_freq: &inv_freq_l,
14995                                rotary_dim: rd_l,
14996                                scale: self.attn_scale,
14997                                softcap: self.attn_softcap,
14998                                window: self.layer_window(li),
14999                                v_norm: self.attn_v_norm,
15000                                qk_norm_after_rope: self.qk_norm_after_rope,
15001                                q_norm: q_norm.as_deref(),
15002                                k_norm: k_norm.as_deref(),
15003                                output_gate: *output_gate,
15004                                softplus_gate: softplus_gate
15005                                    .as_ref()
15006                                    .map(|(gate, per_head)| (gate, *per_head)),
15007                                rope_scale: self.layer_rope_scale(li),
15008                                bias: bias
15009                                    .as_ref()
15010                                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
15011                                rms_eps: eps,
15012                                norm_style: self.norm_style,
15013                                pool: pool.as_deref(),
15014                                v_head_dim: self.layer_v_dim(li),
15015                            };
15016                            attention::qwen_attention(
15017                                &self.ws.n1,
15018                                wq,
15019                                wk,
15020                                wv,
15021                                wo,
15022                                &mut self.kv_cache.layers[li],
15023                                &cfg,
15024                            )
15025                        }
15026                    }
15027                }
15028            };
15029            // Gemma sandwich norm: normalize the attention branch before
15030            // it joins the residual stream.
15031            let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
15032                Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
15033                None => attn_out,
15034            };
15035            let lw = &self.weights.layers[self.phys_layer(li)];
15036            let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
15037            inference::add_rmsnorm_fused_into(
15038                &mut h,
15039                &attn_out,
15040                &lw.post_norm,
15041                self.rms_eps,
15042                self.norm_style,
15043                &mut self.ws.p1,
15044            );
15045            drop(prof);
15046            let mut attn_out = attn_out;
15047            attention::recycle_buf(&mut attn_out);
15048            let post_normed = &self.ws.p1;
15049
15050            let ffn_masked = task_mask
15051                .map(|m| m.ffn_active_count(li) < self.intermediate_size)
15052                .unwrap_or(false);
15053            // One masked dense CONTRACT, dispatched by cost. The
15054            // activation-zeroing arm (the batched sweep's, validated
15055            // against the replica to 0.8%) computes the FULL fused FFN
15056            // and zeroes the dead — right whenever most neurons live.
15057            // The sparse arm reads ONLY active rows and down columns —
15058            // per-row dots are slower per element than the fused kernel,
15059            // so it pays only once the mask is deep enough. The 0.5
15060            // crossover is first-principles (fused kernels run ~2x the
15061            // per-row dot throughput); a shallow specialist (95% alive)
15062            // stays fused, a --target-sparsity bake flips arms on its
15063            // own weight.
15064            let ffn_out = match (ffn_masked, &lw.ffn) {
15065                // A defragged tube layer answers its own mask: the core
15066                // always runs, each tube runs when its bit is on, and
15067                // the tubes that are off are never read from the mmap.
15068                (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
15069                    let row = task_mask
15070                        .and_then(|tm| tm.ffn_masks.get(li))
15071                        .map(|v| v.as_slice());
15072                    tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
15073                }
15074                (true, FfnKind::Dense(d)) => {
15075                    let tm = task_mask.unwrap();
15076                    let alive = tm.ffn_active_count(li);
15077                    let deep = alive * 2 <= self.intermediate_size;
15078                    if deep && d.down_proj.sparse_col_ok() && !d.gate_proj.has_prism_contract() {
15079                        let active = tm.ffn_active_indices(li);
15080                        sparse_ffn_quant(
15081                            d,
15082                            post_normed,
15083                            &active,
15084                            self.hidden_size,
15085                            self.pool.as_deref(),
15086                        )
15087                    } else if deep
15088                        && let (Some(g), Some(u), Some(dn)) = (
15089                            d.gate_proj.as_f32(),
15090                            d.up_proj.as_f32(),
15091                            d.down_proj.as_f32(),
15092                        )
15093                    {
15094                        let active = tm.ffn_active_indices(li);
15095                        inference::sparse_ffn_forward(
15096                            post_normed,
15097                            g,
15098                            u,
15099                            dn,
15100                            self.hidden_size,
15101                            self.intermediate_size,
15102                            &active,
15103                            self.pool.as_deref(),
15104                        )
15105                    } else {
15106                        let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
15107                        dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
15108                    }
15109                }
15110                (true, FfnKind::Moe(m)) => {
15111                    // MoE is sparse by expert selection; a task mask
15112                    // narrows the ROUTABLE set via its expert fields
15113                    // (spec §5) when it carries them.
15114                    let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
15115                    ffn_forward(
15116                        &lw.ffn,
15117                        post_normed,
15118                        self.pool.as_deref(),
15119                        allowed.as_deref(),
15120                    )
15121                }
15122                (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
15123                    dm,
15124                    post_normed,
15125                    &h,
15126                    self.rms_eps,
15127                    self.norm_style,
15128                    self.pool.as_deref(),
15129                ),
15130                (false, _) => match &lw.ffn {
15131                    FfnKind::DenseMoe(dm) => dense_moe_ffn(
15132                        dm,
15133                        post_normed,
15134                        &h,
15135                        self.rms_eps,
15136                        self.norm_style,
15137                        self.pool.as_deref(),
15138                    ),
15139                    FfnKind::Moe(m)
15140                        if task_mask.is_none() && self.mimo_moe.is_dynamic(li, host_tail) =>
15141                    {
15142                        moe_ffn_banked(&mut self.mimo_moe, li, m, post_normed, self.pool.as_deref())
15143                    }
15144                    _ => {
15145                        let allowed = match (&lw.ffn, task_mask) {
15146                            (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
15147                            _ => None,
15148                        };
15149                        ffn_forward(
15150                            &lw.ffn,
15151                            post_normed,
15152                            self.pool.as_deref(),
15153                            allowed.as_deref(),
15154                        )
15155                    }
15156                },
15157            };
15158            let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
15159                Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
15160                None => ffn_out,
15161            };
15162            for (i, &f) in ffn_out.iter().enumerate() {
15163                h[i] += f;
15164            }
15165            let mut ffn_out = ffn_out;
15166            attention::recycle_buf(&mut ffn_out);
15167
15168            // Gemma-4: the layer output is scaled by a learned scalar.
15169            if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
15170                for v in h.iter_mut() {
15171                    *v *= sc;
15172                }
15173            }
15174            // CMF_LAYER_DUMP: this position's hidden after layer li.
15175            if self.layer_dump.is_some() {
15176                self.dump_layer_row(position, li, &h);
15177            }
15178
15179            // Looped Transformer: apply final norm at the end of each loop iteration.
15180            // Nanbeige 4.2: after layer 21 (virtual), apply norm before looping back to layer 0.
15181            if self.is_loop_end(li) && li + 1 < self.num_layers {
15182                h = inference::rms_norm(
15183                    &h,
15184                    &self.weights.final_norm,
15185                    self.rms_eps,
15186                    self.norm_style,
15187                );
15188            }
15189
15190            // Dynamic routing φ capture (on-policy): the
15191            // EMA of the post-residual hidden at the router's phi_layer,
15192            // updated as the context evolves during decode.
15193            if self.dyn_phi_layer == Some(li) {
15194                self.update_dyn_phi(&h);
15195            }
15196        }
15197        crate::gpu::set_layer(-1); // layers done — lm_head outside layer-split
15198        if let Some(t) = t_race_cpu {
15199            crate::gpu::graph_race_record(false, t.elapsed());
15200        }
15201
15202        h
15203    }
15204
15205    /// EMA of φ at the router layer (rolling, weight 0.2 = ~5-token
15206    /// horizon). First observation seeds it exactly.
15207    fn update_dyn_phi(&mut self, h: &[f32]) {
15208        const A: f32 = 0.2;
15209        if self.dyn_phi_ema.len() != h.len() {
15210            self.dyn_phi_ema = vec![0.0; h.len()];
15211            self.dyn_phi_seen = 0;
15212        }
15213        if self.dyn_phi_seen == 0 {
15214            self.dyn_phi_ema.copy_from_slice(h);
15215        } else {
15216            for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
15217                *e = (1.0 - A) * *e + A * v;
15218            }
15219        }
15220        self.dyn_phi_seen += 1;
15221    }
15222
15223    /// Current router φ (EMA at phi_layer); empty until first capture.
15224    pub fn dyn_phi(&self) -> &[f32] {
15225        &self.dyn_phi_ema
15226    }
15227
15228    /// Enable/disable φ capture at the router layer, reset the EMA.
15229    pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
15230        self.dyn_phi_layer = layer;
15231        self.dyn_phi_ema.clear();
15232        self.dyn_phi_seen = 0;
15233    }
15234
15235    /// Skills eligible for dynamic switching: (index, id, phi_layer).
15236    pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
15237        let Some(model) = &self.model else {
15238            return Vec::new();
15239        };
15240        model
15241            .header
15242            .skills
15243            .iter()
15244            .enumerate()
15245            .filter_map(|(i, sk)| {
15246                // A v2 record routes only through the request-level
15247                // backbone-gated decision (its status and gate are
15248                // checked there), never per token.
15249                if sk.is_v2() {
15250                    return None;
15251                }
15252                let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
15253                let sel = sk.selection.as_ref()?;
15254                (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
15255            })
15256            .collect()
15257    }
15258
15259    /// Index of the currently overlaid skill (None = backbone).
15260    pub fn active_skill(&self) -> Option<usize> {
15261        self.dyn_active
15262    }
15263
15264    /// Enable dynamic per-token skill routing: build the hysteresis
15265    /// router from the container's routable skills, start φ capture at
15266    /// their (shared) phi_layer. Returns the number of routable skills
15267    /// (0 = nothing to route; router stays off). Idempotent.
15268    pub fn enable_dynamic_routing(&mut self) -> usize {
15269        use crate::swarm::{DynRouter, RoutableSkill};
15270        let Some(model) = self.model.clone() else {
15271            return 0;
15272        };
15273        // Router policy v2 (spec §9.4) routes per REQUEST: the backbone is
15274        // the default and only the backbone-gated decision may pick a
15275        // skill. A per-token switch would bypass that gate (and change
15276        // the O(1) state mid-sequence), so the hysteresis router never
15277        // runs on such a file; the caller keeps the request-level
15278        // decision.
15279        if let Some(r) = &model.header.router {
15280            tracing::warn!(
15281                "dynamic routing disabled: this file declares router policy '{}' with \
15282                 granularity \"{}\" — the request-level decision applies instead",
15283                r.policy,
15284                r.granularity
15285            );
15286            return 0;
15287        }
15288        // Format-v2 skill records (bit SKILLS_V2) without a router policy:
15289        // their status/gate contract ("auto-routing requires active +
15290        // measured") lives in the backbone-gated decision only — the
15291        // hysteresis router would switch into a quarantined record
15292        // (fail-open). Refuse the whole file, not just its v2 records.
15293        if model.required_features & cortiq_core::format::features::SKILLS_V2 != 0
15294            || model.header.skills.iter().any(|s| s.is_v2())
15295        {
15296            tracing::warn!(
15297                "dynamic routing disabled: this file carries format-v2 skill records \
15298                 (SKILLS_V2) — they route per request through a router policy only"
15299            );
15300            return 0;
15301        }
15302        // A blend materialized f32 working tensors into the layers; there
15303        // is no single skill index to revert from → refuse (honest).
15304        if self.dyn_blend_loaded {
15305            tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
15306            return 0;
15307        }
15308        // A statically-overlaid skill that is NOT FFN-eligible can't be
15309        // cheaply reverted at generation start → refuse rather than
15310        // silently keep it overlaid.
15311        if let Some(a) = self.dyn_active {
15312            if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
15313                tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
15314                return 0;
15315            }
15316        }
15317        let hidden = self.hidden_size;
15318        let mut skills = Vec::new();
15319        for (idx, id, _phi) in self.dynamic_skills() {
15320            if let Some(sel) = model.header.skills[idx].selection.as_ref() {
15321                if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
15322                    skills.push(rs);
15323                }
15324            }
15325        }
15326        if skills.is_empty() {
15327            return 0;
15328        }
15329        // Skills should share a phi_layer; warn (not fail) if they don't.
15330        let phi = skills[0].phi_layer;
15331        if skills.iter().any(|s| s.phi_layer != phi) {
15332            tracing::warn!("routable skills disagree on phi_layer; using {phi}");
15333        }
15334        let n = skills.len();
15335        self.set_dyn_phi_layer(Some(phi));
15336        self.dyn_router = Some(DynRouter::new(skills));
15337        n
15338    }
15339
15340    /// Human-readable switch log from the last dynamic-routed generation.
15341    pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
15342        self.dyn_router
15343            .as_ref()
15344            .map(|r| r.switches.clone())
15345            .unwrap_or_default()
15346    }
15347
15348    /// LM head: hidden → logits [vocab_size]. The dominant matvec of
15349    /// every decode step — row-parallel on the worker pool.
15350    fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
15351        let _mimo_q8 = self.mimo_moe.is_on()
15352            .then(crate::qtensor::enter_full_gpu_q8_scope);
15353        let rows = self.weights.lm_head.rows();
15354        let mut logits = attention::take_buf(rows.min(self.vocab_size));
15355        // Banked MiMo uses the same exact projection family for the
15356        // plain/draft head and the batched verification head. Read both
15357        // scale planes in-place instead of preparing per-op scale buffers.
15358        let served = self.mimo_moe.is_on() && crate::gpu::mimo_q8_short_enabled()
15359            && rows == self.vocab_size && !self.weights.lm_head.has_prism_contract()
15360            && self.weights.lm_head.graph_weight().is_some_and(|(model, idx, kind, _)| {
15361                kind == 7 && crate::gpu::q82_short_rows(model, idx, hidden, 1,
15362                    rows, self.hidden_size, &mut logits)
15363            });
15364        if !served {
15365            self.weights.lm_head.matvec(hidden, &mut logits, self.pool.as_deref());
15366        }
15367        logits.resize(self.vocab_size, 0.0);
15368        if let Some(m) = self.logit_multiplier {
15369            for l in logits.iter_mut() {
15370                *l *= m;
15371            }
15372        }
15373        if let Some(c) = self.final_softcap {
15374            for l in logits.iter_mut() {
15375                *l = c * (*l / c).tanh();
15376            }
15377        }
15378        if let Some(cm) = self.head_clusters.as_ref() {
15379            self.hierarchical_head_logprobs(hidden, cm, &mut logits);
15380        }
15381        logits
15382    }
15383
15384    /// Two-level head (Cortiq Embryo): in place, logits[v] ← log p(v) =
15385    /// (lc[c] − lse(lc)) + (logit[v] − lse over v's cluster block), c = v / S.
15386    fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
15387        let h = hidden.len();
15388        let ncl = cm.len() / h.max(1);
15389        if ncl == 0 || logits.len() % ncl != 0 {
15390            return;
15391        }
15392        let cs = logits.len() / ncl;
15393        // cluster logits + log-softmax
15394        let mut lc = vec![0.0f32; ncl];
15395        for c in 0..ncl {
15396            let row = &cm[c * h..(c + 1) * h];
15397            let mut s = 0.0f32;
15398            for j in 0..h {
15399                s += row[j] * hidden[j];
15400            }
15401            lc[c] = s;
15402        }
15403        let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15404        let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
15405        for c in 0..ncl {
15406            let blk = &mut logits[c * cs..(c + 1) * cs];
15407            let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15408            let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
15409            let add = lc[c] - lse - bl;
15410            for v in blk.iter_mut() {
15411                *v += add;
15412            }
15413        }
15414    }
15415
15416    /// Prefill `ids` and return the next-token logits — what the model
15417    /// would predict next, WITHOUT committing to generation (introspection
15418    /// for `cortiq explain`). Clears and repopulates the KV cache; leaves
15419    /// the active overlay untouched.
15420    pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
15421        #[cfg(target_os = "macos")]
15422        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
15423        self.clear_sequence_state();
15424        // This helper is used by the pooled classification endpoint, where
15425        // every request is a fresh sequence. The shared reset also clears the
15426        // wgpu token graph's device-side recurrent state.
15427        crate::gpu::graph_race_begin_generation();
15428        if task_mask.is_none() {
15429            self.o1_begin();
15430        }
15431        let mut hidden = vec![0.0f32; self.hidden_size];
15432        for (pos, &id) in ids.iter().enumerate() {
15433            let emb = self.embed_single(id);
15434            hidden = self.forward_layers(&emb, pos, task_mask);
15435        }
15436        if let Err(err) = self.o1_seal_checked() {
15437            self.o1_fail(err);
15438        }
15439        // Stacks that own their head (V4, V4.1, Qwen3.8-Flash-Next, GLM-5)
15440        // return a zero hidden and hand the logits out of band.
15441        if let Some(logits) = self.graph_logits.take() {
15442            return logits;
15443        }
15444        inference::rms_norm_into(
15445            &hidden,
15446            &self.weights.final_norm,
15447            self.rms_eps,
15448            self.norm_style,
15449            &mut self.ws.n1,
15450        );
15451        self.lm_head_forward(&self.ws.n1)
15452    }
15453}
15454
15455/// Convenience: deterministic tiny pipeline for tests.
15456pub fn create_test_pipeline(
15457    hidden_size: usize,
15458    intermediate_size: usize,
15459    num_heads: usize,
15460    num_kv_heads: usize,
15461    head_dim: usize,
15462    num_layers: usize,
15463    vocab_size: usize,
15464) -> Pipeline {
15465    // Small pseudo-random weights: constant weights make attention
15466    // degenerate and hide indexing bugs.
15467    let synth = |n: usize, salt: usize| -> Vec<f32> {
15468        (0..n)
15469            .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
15470            .collect()
15471    };
15472    let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
15473        QTensor::from_f32(synth(rows * cols, salt), rows, cols)
15474    };
15475    let layer_weights: Vec<LayerWeights> = (0..num_layers)
15476        .map(|li| LayerWeights {
15477            input_norm: vec![1.0; hidden_size],
15478            post_norm: vec![1.0; hidden_size],
15479            attn_out_norm: None,
15480            ffn_out_norm: None,
15481            layer_scale: None,
15482            ffn: FfnKind::Dense(DenseFfn {
15483                gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
15484                up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
15485                down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
15486                act: Act::Silu,
15487                down_t: None,
15488                segs: Vec::new(),
15489            }),
15490            attn: AttnKind::Full {
15491                bias: None,
15492                wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
15493                wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
15494                wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
15495                wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
15496                q_norm: None,
15497                k_norm: None,
15498                output_gate: false,
15499                softplus_gate: None,
15500            },
15501        })
15502        .collect();
15503
15504    Pipeline::new(
15505        Tokenizer::byte_level(),
15506        PipelineWeights {
15507            embed_tokens: qt(vocab_size, hidden_size, 100),
15508            layers: layer_weights,
15509            lm_head: qt(vocab_size, hidden_size, 200),
15510            final_norm: vec![1.0; hidden_size],
15511        },
15512        hidden_size,
15513        intermediate_size,
15514        num_heads,
15515        num_kv_heads,
15516        head_dim,
15517        num_layers,
15518        num_layers, // physical_layers = num_layers (non-looped)
15519        false,      // loop_final_norm
15520        vocab_size,
15521        1e-6,
15522        10_000.0,
15523        NormStyle::Qwen,
15524        4096,
15525        SamplerConfig {
15526            seed: Some(42),
15527            ..Default::default()
15528        },
15529    )
15530}
15531
15532/// Batched dense-FFN: gate/up/down via matmat (element-wise the same
15533/// math as b × dense_ffn — the same dot kernels).
15534/// One mask bit, LSB-first per byte — `TaskMask::ffn_active_indices`'s
15535/// convention.
15536#[inline]
15537fn mask_bit(row: &[u8], j: usize) -> bool {
15538    (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
15539}
15540
15541/// Zero the CLOSED neurons' activations in a [rows × inter] panel — the
15542/// masked-inference fast path's whole trick: full fused quant compute,
15543/// then the mask lands on the ACTIVATIONS, which is arithmetically the
15544/// pruned network without touching a quantized weight byte. Whole open
15545/// bytes (0xFF = 8 open neurons) skip in one test.
15546/// `CMF_FFN_MASK_GAIN` — Patent 12 FIG. 4, variance-preserving
15547/// rescaling: truncation removes a share of the layer's output energy,
15548/// so the survivors are scaled up to put the variance back where the
15549/// downstream norm expects it. A scalar here; per layer it is
15550/// `sqrt(total energy / kept energy)`.
15551fn mask_gain() -> f32 {
15552    static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
15553    *G.get_or_init(|| {
15554        std::env::var("CMF_FFN_MASK_GAIN")
15555            .ok()
15556            .and_then(|v| v.parse().ok())
15557            .unwrap_or(1.0)
15558    })
15559}
15560
15561fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
15562    // With CMF_FFN_MEANFILL a closed neuron contributes its average
15563    // instead of nothing — same bytes read, one constant restored.
15564    let fill = meanfill().and_then(|(i, v)| {
15565        let li = crate::gpu::cur_layer();
15566        (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
15567    });
15568    for r in 0..rows {
15569        let base = r * inter;
15570        for (bi, &byte) in row.iter().enumerate() {
15571            if byte == 0xFF {
15572                continue;
15573            }
15574            let j0 = bi * 8;
15575            for bit in 0..8 {
15576                let j = j0 + bit;
15577                if j < inter && byte & (1 << bit) == 0 {
15578                    g[base + j] = fill.map_or(0.0, |f| f[j]);
15579                }
15580            }
15581        }
15582    }
15583    let gain = mask_gain();
15584    if gain != 1.0 {
15585        for v in g[..rows * inter].iter_mut() {
15586            *v *= gain;
15587        }
15588    }
15589}
15590
15591/// True when neuron `i`'s bit is set (no mask = everything runs).
15592#[inline]
15593fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
15594    row.is_none_or(|r| mask_bit(r, i))
15595}
15596
15597/// Every bit below `n` set — the common case for a tube file's CORE,
15598/// where only the tube bits vary per task.
15599fn all_bits_on(row: &[u8], n: usize) -> bool {
15600    (0..n).all(|i| mask_bit(row, i))
15601}
15602
15603/// `CMF_TUBE_TOPK` — how many tubes a TOKEN may open (0 = the task mask
15604/// decides alone). This is the dense FFN read as a mixture: the tubes
15605/// are the experts a k-means over `gate_proj` rows found, and the token
15606/// picks among them. `CMF_TUBE_SCORE=gate` scores a tube by its own
15607/// gate (realizable: only `up`/`down` of the losers go unread),
15608/// `=oracle` scores by the true `silu(gate)·up` mass (the ceiling —
15609/// only `down` is saved, and the selection has read what it predicts).
15610fn tube_topk() -> usize {
15611    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
15612    *K.get_or_init(|| {
15613        std::env::var("CMF_TUBE_TOPK")
15614            .ok()
15615            .and_then(|v| v.parse().ok())
15616            .unwrap_or(0)
15617    })
15618}
15619
15620fn tube_score_oracle() -> bool {
15621    static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15622    *O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
15623}
15624
15625/// The routed arm of `tube_ffn`: a token opens only its best `k` tubes.
15626/// At `b == 1` (decode) the losers are genuinely never read — that is
15627/// the speed. At `b > 1` (the scoring sweep) every tube is computed and
15628/// the losers' activations are zeroed instead: same arithmetic, so the
15629/// perplexity is the routed model's, measured without a per-token
15630/// gather in the middle of a GEMM.
15631fn tube_ffn_routed(
15632    d: &DenseFfn,
15633    xs: &[f32],
15634    b: usize,
15635    pool: Option<&Pool>,
15636    mask_row: Option<&[u8]>,
15637    k: usize,
15638) -> Vec<f32> {
15639    let hidden = d.down_proj.rows();
15640    let core = d.gate_proj.rows();
15641    let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
15642    let mut out = match (b, core_full, mask_row) {
15643        (1, true, _) => dense_ffn(d, xs, pool),
15644        (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
15645        (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
15646        (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
15647    };
15648    let cand: Vec<usize> = (0..d.segs.len())
15649        .filter(|&i| tube_bit(mask_row, d.segs[i].start))
15650        .collect();
15651    if cand.is_empty() {
15652        return out;
15653    }
15654    // gate (and, where the score or the batch needs it, up) per tube.
15655    // The SCORE is taken at the point the serving path could take it:
15656    // off the gate alone, or off the finished activation for the oracle.
15657    let oracle = tube_score_oracle();
15658    let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
15659    let mut scores = vec![0f32; b * cand.len()];
15660    for (ci, &i) in cand.iter().enumerate() {
15661        let seg = &d.segs[i];
15662        let w = seg.width;
15663        let mut g = vec![0.0f32; b * w];
15664        if b == 1 {
15665            seg.gate.matvec(xs, &mut g, pool);
15666        } else {
15667            seg.gate.matmat(xs, b, &mut g, pool);
15668        }
15669        for v in g.iter_mut() {
15670            *v = Act::Silu.combine(*v, 1.0);
15671        }
15672        if !oracle {
15673            for t in 0..b {
15674                scores[t * cand.len() + ci] =
15675                    g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15676            }
15677        }
15678        if oracle || b > 1 {
15679            let mut u = vec![0.0f32; b * w];
15680            if b == 1 {
15681                seg.up.matvec(xs, &mut u, pool);
15682            } else {
15683                seg.up.matmat(xs, b, &mut u, pool);
15684            }
15685            for (a, &v) in g.iter_mut().zip(u.iter()) {
15686                *a *= v;
15687            }
15688            if oracle {
15689                for t in 0..b {
15690                    scores[t * cand.len() + ci] =
15691                        g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
15692                }
15693            }
15694        }
15695        acts.push(g);
15696    }
15697    // per-token scores and the winners
15698    let keep = k.min(cand.len());
15699    let mut scratch: Vec<f32> = Vec::new();
15700    for t in 0..b {
15701        let mut sc: Vec<(f32, usize)> = (0..cand.len())
15702            .map(|ci| (scores[t * cand.len() + ci], ci))
15703            .collect();
15704        sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
15705        let mut alive = vec![false; cand.len()];
15706        for &(_, ci) in sc.iter().take(keep) {
15707            alive[ci] = true;
15708        }
15709        if b > 1 {
15710            for (ci, a) in acts.iter_mut().enumerate() {
15711                if !alive[ci] {
15712                    let w = d.segs[cand[ci]].width;
15713                    a[t * w..(t + 1) * w].fill(0.0);
15714                }
15715            }
15716        } else {
15717            // decode: finish only the winners — the losers' up/down
15718            // (and, with the gate score, everything but their gate)
15719            // are never touched.
15720            for (ci, &i) in cand.iter().enumerate() {
15721                if !alive[ci] {
15722                    continue;
15723                }
15724                let seg = &d.segs[i];
15725                let w = seg.width;
15726                let g = &mut acts[ci];
15727                if !tube_score_oracle() {
15728                    scratch.clear();
15729                    scratch.resize(w, 0.0);
15730                    seg.up.matvec(xs, &mut scratch, pool);
15731                    for (a, &v) in g.iter_mut().zip(scratch.iter()) {
15732                        *a *= v;
15733                    }
15734                }
15735                let mut acc = vec![0.0f32; hidden];
15736                seg.down.matvec(g, &mut acc, pool);
15737                for (o, a) in out.iter_mut().zip(&acc) {
15738                    *o += *a;
15739                }
15740            }
15741        }
15742    }
15743    if b > 1 {
15744        for (ci, &i) in cand.iter().enumerate() {
15745            let seg = &d.segs[i];
15746            let mut acc = vec![0.0f32; b * hidden];
15747            seg.down.matmat(&acts[ci], b, &mut acc, pool);
15748            for (o, a) in out.iter_mut().zip(&acc) {
15749                *o += *a;
15750            }
15751        }
15752    }
15753    out
15754}
15755
15756/// FFN of a defragged tube layer: the always-on core plus the tubes the
15757/// task mask switches on. Each tube is a normal tensor triple, so the
15758/// same kernels run it and an inactive tube's bytes are never read —
15759/// that is the whole point of the defrag (a scattered mask cannot skip
15760/// bytes; a contiguous one is just a smaller matrix).
15761fn tube_ffn(
15762    d: &DenseFfn,
15763    xs: &[f32],
15764    b: usize,
15765    pool: Option<&Pool>,
15766    mask_row: Option<&[u8]>,
15767) -> Vec<f32> {
15768    if tube_topk() > 0 {
15769        return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
15770    }
15771    let hidden = d.down_proj.rows();
15772    let core = d.gate_proj.rows();
15773    let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
15774    let mut out = match (b, core_full, mask_row) {
15775        (1, true, _) => dense_ffn(d, xs, pool),
15776        (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
15777        (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
15778        (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
15779    };
15780    TUBE_SCRATCH.with(|sc| {
15781        let mut sc = sc.borrow_mut();
15782        let [g, u, acc] = &mut *sc;
15783        for seg in &d.segs {
15784            if !tube_bit(mask_row, seg.start) {
15785                continue;
15786            }
15787            let w = seg.width;
15788            g.resize(b * w, 0.0);
15789            if b == 1
15790                && d.act == Act::Silu
15791                && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
15792            {
15793                // g holds silu(gate)·up.
15794            } else {
15795                u.resize(b * w, 0.0);
15796                if b == 1 {
15797                    QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
15798                } else {
15799                    seg.gate.matmat(xs, b, g, pool);
15800                    seg.up.matmat(xs, b, u, pool);
15801                }
15802                for i in 0..b * w {
15803                    g[i] = d.act.combine(g[i], u[i]);
15804                }
15805            }
15806            acc.resize(b * hidden, 0.0);
15807            acc.fill(0.0);
15808            if b == 1 {
15809                seg.down.matvec(g, acc, pool);
15810            } else {
15811                seg.down.matmat(g, b, acc, pool);
15812            }
15813            for (o, a) in out.iter_mut().zip(acc.iter()) {
15814                *o += *a;
15815            }
15816        }
15817        out
15818    })
15819}
15820
15821thread_local! {
15822    /// gate / up / down-accumulator scratch for the tube loop — a tube
15823    /// runs once per layer per token, and a fresh Vec each time is a
15824    /// malloc per tube per layer per token.
15825    static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
15826        const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
15827}
15828
15829fn dense_ffn_batch(
15830    d: &DenseFfn,
15831    xs: &[f32],
15832    b: usize,
15833    pool: Option<&Pool>,
15834    mask_row: Option<&[u8]>,
15835) -> Vec<f32> {
15836    let inter = d.gate_proj.rows();
15837    let hidden = d.down_proj.rows();
15838    // Fused on-device SwiGLU when the device is in play: three separate
15839    // `matmat` calls are three round trips per layer, and the gate/up
15840    // panels (b × inter — 22 MB each at a 512-token chunk) cross the bus
15841    // twice for nothing. The kernel already existed for the image DiT;
15842    // the LLM prefill was simply never wired to it. A task mask needs the
15843    // activations on the host between the halves, so it keeps the CPU
15844    // arm below.
15845    if mask_row.is_none()
15846        && d.act == Act::Silu
15847        && b >= 32
15848        && crate::gpu::enabled_here()
15849        && !crate::gpu::mm_killed()
15850        // The refit pass needs this layer's activations on the host; the
15851        // fused chain keeps them on the device. Refusing it here costs
15852        // one round trip and keeps every GEMM on the card — the
15853        // alternative was running the whole calibration on the CPU.
15854        && refit_dir().is_none()
15855        // Same for the mass/hit probes. The accumulator at the bottom of
15856        // this function only sees `g` when `g` came back to the host, so
15857        // a fused batch would leave it summing nothing — a probe that
15858        // reports zeros rather than failing, which is worse.
15859        && !ffn_probe_active()
15860    {
15861        if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
15862            d.gate_proj.mapped_q4t(),
15863            d.up_proj.mapped_q4t(),
15864            d.down_proj.mapped_q4t(),
15865        ) {
15866            let mut out = vec![0.0f32; b * hidden];
15867            if crate::gpu::q4t_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
15868                return out;
15869            }
15870        }
15871        // The q4tp twin (same kernel family, scale from the row ladder) —
15872        // the DiT has run it in production since the pipeline containers;
15873        // the LLM prefill was simply never wired to it, so a q4tp model's
15874        // prefill panels stayed on the CPU.
15875        if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
15876            d.gate_proj.mapped_q4tp(),
15877            d.up_proj.mapped_q4tp(),
15878            d.down_proj.mapped_q4tp(),
15879        ) {
15880            let mut out = vec![0.0f32; b * hidden];
15881            if crate::gpu::q4tp_ffn(model, w1, w3, w2, xs, b, hidden, inter, &mut out) {
15882                return out;
15883            }
15884        }
15885    }
15886    let mut g = vec![0.0f32; b * inter];
15887    d.gate_proj.matmat(xs, b, &mut g, pool);
15888    let mut u = vec![0.0f32; b * inter];
15889    d.up_proj.matmat(xs, b, &mut u, pool);
15890    if gate_topk() > 0 && d.act == Act::Silu {
15891        for t in 0..b {
15892            let row = &mut g[t * inter..(t + 1) * inter];
15893            for v in row.iter_mut() {
15894                *v = Act::Silu.combine(*v, 1.0);
15895            }
15896            keep_top_k(row, gate_topk());
15897        }
15898        for i in 0..b * inter {
15899            g[i] *= u[i];
15900        }
15901    } else {
15902        for i in 0..b * inter {
15903            g[i] = d.act.combine(g[i], u[i]);
15904        }
15905    }
15906    if let Some(row) = mask_row {
15907        zero_masked_cols(&mut g, b, inter, row);
15908    }
15909    if oracle_topk() > 0 {
15910        for t in 0..b {
15911            keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
15912        }
15913    }
15914    let mut out = vec![0.0f32; b * hidden];
15915    d.down_proj.matmat(&g, b, &mut out, pool);
15916    if refit_dir().is_some() {
15917        let li = crate::gpu::cur_layer();
15918        if li >= 0 {
15919            refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
15920        }
15921    }
15922    // The DTG-MA probe, on the batched path: one prefill sweep gives the
15923    // same per-neuron statistic the per-position probe does, and on a 27B
15924    // that is minutes instead of hours.
15925    FFN_PROBE.with(|pr| {
15926        if let Some(acc) = pr.borrow_mut().as_mut() {
15927            let li = crate::gpu::cur_layer();
15928            if li < 0 {
15929                return;
15930            }
15931            let Some(row) = acc.get_mut(li as usize) else {
15932                return;
15933            };
15934            let sq = probe_sq();
15935            for t in 0..b {
15936                for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
15937                    *a += if sq {
15938                        (v as f64) * (v as f64)
15939                    } else {
15940                        (v as f64).abs()
15941                    };
15942                }
15943            }
15944        }
15945    });
15946    out
15947}
15948
15949/// Batched MoE-FFN: router batched, positions are GROUPED by expert —
15950/// an expert's weights are read once for all its positions in the chunk
15951/// (the main prefill-GEMM win on MoE: 960MB/token of 35B experts).
15952/// Accumulate per-channel activation energy for `CMF_RMS_TRACE`.
15953fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
15954    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15955    static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
15956    let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
15957    let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
15958    if (!on && !dump) || b == 0 {
15959        return;
15960    }
15961    let hidden = xs.len() / b;
15962    if on {
15963        let mut acc = m.act_sq.borrow_mut();
15964        if acc.len() < hidden {
15965            acc.resize(hidden, 0.0);
15966        }
15967        for t in 0..b {
15968            let row = &xs[t * hidden..(t + 1) * hidden];
15969            for (a, &v) in acc.iter_mut().zip(row) {
15970                *a += (v as f64) * (v as f64);
15971            }
15972        }
15973    }
15974    if dump {
15975        // Cap the capture: the covariance needs a few thousand rows, and a
15976        // whole prefill of every layer would be gigabytes for no extra rank.
15977        let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
15978            .ok()
15979            .and_then(|v| v.parse().ok())
15980            .unwrap_or(4096);
15981        let mut rows = m.act_rows.borrow_mut();
15982        if rows.len() < cap * hidden {
15983            let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
15984            rows.extend_from_slice(&xs[..take * hidden]);
15985        }
15986    }
15987}
15988
15989/// Send-able cursor over a Vec-of-Vecs: each pool worker writes only its
15990/// own slots (disjoint by construction in the caller).
15991#[derive(Clone, Copy)]
15992struct SendVecs(*mut Vec<f32>);
15993unsafe impl Send for SendVecs {}
15994unsafe impl Sync for SendVecs {}
15995impl SendVecs {
15996    #[inline]
15997    fn at(self, i: usize) -> *mut Vec<f32> {
15998        unsafe { self.0.add(i) }
15999    }
16000}
16001
16002fn moe_ffn_batch(
16003    m: &MoeFfn,
16004    xs: &[f32],
16005    b: usize,
16006    hidden: usize,
16007    pool: Option<&Pool>,
16008    allowed: Option<&[bool]>,
16009) -> Vec<f32> {
16010    accumulate_act(m, xs, b);
16011    let ne = m.experts.len();
16012    let mut logits = vec![0.0f32; b * ne];
16013    match &m.resonance {
16014        Some(r) => {
16015            let hdim = xs.len() / b.max(1);
16016            for bi in 0..b {
16017                r.scores(
16018                    &xs[bi * hdim..(bi + 1) * hdim],
16019                    &mut logits[bi * ne..(bi + 1) * ne],
16020                );
16021            }
16022        }
16023        None => m.router.matmat(xs, b, &mut logits, pool),
16024    }
16025
16026    // Assignments: expert → [(position, weight)] — same routing as
16027    // moe_ffn, per position (see `moe_route`).
16028    let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
16029    {
16030        let mut st = m.stats.borrow_mut();
16031        if st.len() < ne {
16032            st.resize(ne, 0);
16033        }
16034        for bi in 0..b {
16035            let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
16036            for &e in &idx {
16037                st[e] += 1;
16038                assign[e].push((bi, p[e] / wsum));
16039            }
16040        }
16041    }
16042
16043    let mut out = vec![0.0f32; b * hidden];
16044    let cols = m.experts[0].gate_proj.cols();
16045    let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
16046        let sb = list.len();
16047        let mut sub = vec![0.0f32; sb * cols];
16048        for (k, &(bi, _)) in list.iter().enumerate() {
16049            sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16050        }
16051        let eo = dense_ffn_batch(d, &sub, sb, pool, None);
16052        for (k, &(bi, w)) in list.iter().enumerate() {
16053            for i in 0..hidden {
16054                out[bi * hidden + i] += w * eo[k * hidden + i];
16055            }
16056        }
16057    };
16058    // Routed experts: the panels are TINY (b·top_k spread over every
16059    // expert — a few positions each), so a pool dispatch per expert is
16060    // pure barrier cost. Invert the parallelism: workers take WHOLE
16061    // experts (serial math inside), then one deterministic scatter in
16062    // expert order — the exact accumulation order the serial loop had.
16063    let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
16064    if pool.is_some() && active.len() >= 8 {
16065        let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
16066        {
16067            let panel_ptr = SendVecs(panels.as_mut_ptr());
16068            // Capture only the expert table: `m` itself carries RefCell
16069            // stats and must not cross the pool boundary.
16070            let experts = &m.experts;
16071            let (active_r, assign_r) = (&active, &assign);
16072            let inherit_cpu = crate::gpu::inherit_cpu_scope();
16073            let run = |start: usize, end: usize| {
16074                let _cpu_scope = inherit_cpu();
16075                for ai in start..end {
16076                    let e = active_r[ai];
16077                    let list = &assign_r[e];
16078                    let sb = list.len();
16079                    let mut sub = vec![0.0f32; sb * cols];
16080                    for (k, &(bi, _)) in list.iter().enumerate() {
16081                        sub[k * cols..(k + 1) * cols]
16082                            .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16083                    }
16084                    // SAFETY: each worker owns a disjoint panels[ai].
16085                    unsafe {
16086                        *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
16087                    }
16088                }
16089            };
16090            match pool {
16091                Some(p) => p.run_rows(active.len(), &run),
16092                None => run(0, active.len()),
16093            }
16094        }
16095        for (ai, &e) in active.iter().enumerate() {
16096            for (k, &(bi, w)) in assign[e].iter().enumerate() {
16097                let eo = &panels[ai][k * hidden..(k + 1) * hidden];
16098                for i in 0..hidden {
16099                    out[bi * hidden + i] += w * eo[i];
16100                }
16101            }
16102        }
16103    } else {
16104        for &e in &active {
16105            run_expert(&m.experts[e], &assign[e], &mut out);
16106        }
16107    }
16108    if let Some((se, gate)) = &m.shared {
16109        let all: Vec<(usize, f32)> = if let Some(gate) = gate {
16110            let mut gl = vec![0.0f32; b];
16111            gate.matmat(xs, b, &mut gl, pool);
16112            (0..b)
16113                .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
16114                .collect()
16115        } else {
16116            (0..b).map(|bi| (bi, 1.0)).collect()
16117        };
16118        run_expert(se, &all, &mut out);
16119    }
16120    out
16121}
16122
16123/// Decode-exact multi-token MoE — the MiMo speculative verify's FFN. Row
16124/// `r` of the result is bit-identical to `moe_ffn(m, x_r)` on the CPU
16125/// (`moe_ffn_cpu` → `moe_ffn_cpu_batched`): router matvec per row, the same
16126/// routing, the same int8 gate/up/SiLU and down terms
16127/// (`QTensor::moe_gate_up_rows` / `moe_down_rows`), and the row's experts
16128/// summed in ITS route order from 0. What the rows share is the weight
16129/// traffic: each routed expert is read once for every row that picked it.
16130/// (`moe_ffn_batch`, the prompt path, groups the same way but sums in
16131/// expert-index order and runs blocked kernels on wide groups — close, not
16132/// bit-equal to decode.) Any layer the kernels do not cover, or a device
16133/// that could answer `moe_ffn` itself, walks `moe_ffn` row by row.
16134fn moe_ffn_rows_exact(
16135    m: &MoeFfn,
16136    xs: &[f32],
16137    b: usize,
16138    hidden: usize,
16139    pool: Option<&Pool>,
16140) -> Vec<f32> {
16141    let mut out = vec![0.0f32; b * hidden];
16142    let per_row = |out: &mut [f32]| {
16143        for r in 0..b {
16144            let o = moe_ffn(m, &xs[r * hidden..(r + 1) * hidden], pool, None);
16145            out[r * hidden..(r + 1) * hidden].copy_from_slice(&o);
16146        }
16147    };
16148    let covered = !crate::gpu::enabled_here()
16149        && moe_batch_enabled()
16150        && m.shared.is_none()
16151        && m.resonance.is_none()
16152        && FFN_PROBE.with(|pr| pr.borrow().is_none())
16153        && m.experts.iter().all(|d| d.act == Act::Silu);
16154    if !covered {
16155        per_row(&mut out);
16156        return out;
16157    }
16158    let ne = m.experts.len();
16159    // Routing, row by row, exactly as `moe_ffn`.
16160    let mut routes: Vec<(Vec<usize>, Vec<f32>)> = Vec::with_capacity(b);
16161    for r in 0..b {
16162        let x = &xs[r * hidden..(r + 1) * hidden];
16163        accumulate_act(m, x, 1);
16164        let mut logits = vec![0.0f32; ne];
16165        m.router.matvec(x, &mut logits, pool);
16166        let (idx, p, wsum) = moe_route(&logits, m, None);
16167        {
16168            let mut st = m.stats.borrow_mut();
16169            if st.len() < ne {
16170                st.resize(ne, 0);
16171            }
16172            for &e in &idx {
16173                st[e] += 1;
16174            }
16175        }
16176        let w: Vec<f32> = idx
16177            .iter()
16178            .map(|&e| p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]))
16179            .collect();
16180        routes.push((idx, w));
16181    }
16182    if routes.iter().any(|(idx, _)| idx.is_empty()) {
16183        per_row(&mut out);
16184        return out;
16185    }
16186    // Group the (row, expert) picks by expert, in first-seen order.
16187    let mut experts: Vec<usize> = Vec::new();
16188    let mut groups: Vec<Vec<usize>> = Vec::new();
16189    for (r, (idx, _)) in routes.iter().enumerate() {
16190        for &e in idx {
16191            match experts.iter().position(|&x| x == e) {
16192                Some(g) => groups[g].push(r),
16193                None => {
16194                    experts.push(e);
16195                    groups.push(vec![r]);
16196                }
16197            }
16198        }
16199    }
16200    let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
16201    let inter = m.experts[experts[0]].gate_proj.rows();
16202    let pairs: Vec<(&QTensor, &QTensor)> = experts
16203        .iter()
16204        .map(|&e| (&m.experts[e].gate_proj, &m.experts[e].up_proj))
16205        .collect();
16206    let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
16207    if !QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut gs, pool) {
16208        per_row(&mut out);
16209        return out;
16210    }
16211    let downs: Vec<&QTensor> = experts.iter().map(|&e| &m.experts[e].down_proj).collect();
16212    let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
16213    let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; hidden]).collect();
16214    if !QTensor::moe_down_rows(&downs, &lens, &gs, &mut ds, pool) {
16215        per_row(&mut out);
16216        return out;
16217    }
16218    // Where each (row, expert) term landed in the flat pair list.
16219    let mut slot = std::collections::HashMap::with_capacity(n_pairs);
16220    let mut p = 0usize;
16221    for (g, &e) in experts.iter().enumerate() {
16222        for &r in &groups[g] {
16223            slot.insert((r, e), p);
16224            p += 1;
16225        }
16226    }
16227    for (r, (idx, w)) in routes.iter().enumerate() {
16228        let terms: Vec<(&[f32], f32)> = idx
16229            .iter()
16230            .zip(w)
16231            .map(|(&e, &we)| (ds[slot[&(r, e)]].as_slice(), we))
16232            .collect();
16233        let row = &mut out[r * hidden..(r + 1) * hidden];
16234        for (i, dst) in row.iter_mut().enumerate() {
16235            // `moe_down_many`'s per-row sum: from 0, in route order.
16236            let mut acc = 0f32;
16237            for (d, we) in &terms {
16238                acc += we * d[i];
16239            }
16240            *dst = acc;
16241        }
16242    }
16243    out
16244}
16245
16246thread_local! {
16247    /// gate/up activation scratch for the dense FFN paths (single uses
16248    /// two slots, the fused pair all four) — these were fresh
16249    /// intermediate-size Vecs on every layer of every token.
16250    static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
16251        const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
16252}
16253
16254/// Dense SwiGLU FFN through QTensor matvecs (any storage).
16255fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16256    // Per-token sparsity, when the file was built for it: gate first,
16257    // then only the chosen neurons' up/down rows leave the mmap.
16258    if gate_topk() > 0
16259        && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
16260    {
16261        return out;
16262    }
16263    // Whole-FFN GPU submit (этап 4.2 increment): gate → silu·up → down
16264    // chained in ONE command buffer with the intermediate activations
16265    // resident on the device — 3 per-op polls become 1 per layer. The
16266    // moe_block backend already implements exactly this chain; a dense
16267    // FFN is one expert with weight 1. Runtime probe: the chain still
16268    // pays one submit+poll per layer — alternate it against the pure-CPU
16269    // FFN and keep whichever is faster on this machine.
16270    // q1 FFNs offload at any practical size: the q1 CPU kernel is
16271    // compute-bound, so the UMA threshold logic does not apply — the
16272    // probe measures and decides either way.
16273    // The fused GPU block has no descriptor-aware Prism path: it would either
16274    // consume an unrotated activation or decline after inspecting the mixed
16275    // q2tp/q4tp tensors.  Do not let that structural refusal enter the FFN
16276    // probe's CPU_ONLY scope; the ordinary body below dispatches each matrix
16277    // through QTensor::matvec, which owns the signed FWHT + affine q2tp route.
16278    let prism_body = d.gate_proj.has_prism_contract()
16279        || d.up_proj.has_prism_contract()
16280        || d.down_proj.has_prism_contract();
16281    if !prism_body
16282        && crate::gpu::enabled_here()
16283        && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
16284    {
16285        let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
16286            crate::gpu::ProbeArm::Gpu
16287        } else {
16288            crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
16289        };
16290        match arm {
16291            crate::gpu::ProbeArm::Gpu => {
16292                let t0 = std::time::Instant::now();
16293                if let Some(out) = dense_ffn_gpu(d, x, pool) {
16294                    crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
16295                    return out;
16296                }
16297                // Declined: no timing exists, so say so. Silence here is
16298                // what left `ffn` undecided for 9000 calls and cost a
16299                // failed device attempt on half of them.
16300                crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
16301            }
16302            crate::gpu::ProbeArm::CpuTimed => {
16303                let t0 = std::time::Instant::now();
16304                let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16305                crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
16306                return out;
16307            }
16308            crate::gpu::ProbeArm::Cpu => {
16309                return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16310            }
16311        }
16312    }
16313    dense_ffn_cpu(d, x, pool)
16314}
16315
16316/// The pure-CPU dense-FFN body (also the fallback of every GPU refusal).
16317fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16318    let inter = d.gate_proj.rows();
16319    FFN_SCRATCH.with(|s| {
16320        let mut s = s.borrow_mut();
16321        let [g, u, ..] = &mut *s;
16322        g.resize(inter, 0.0);
16323        // Fused gate+up+silu: one dispatch, no separate silu pass.
16324        // Falls back to matvec_many + silu loop for unsupported dtypes.
16325        if gate_topk() > 0 {
16326            // Gate first, select, and only then pay for `up`: the
16327            // measurement arm computes both and zeroes the losers, which
16328            // is the same arithmetic.
16329            u.resize(inter, 0.0);
16330            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16331            for i in 0..inter {
16332                g[i] = Act::Silu.combine(g[i], 1.0);
16333            }
16334            keep_top_k(g, gate_topk());
16335            for i in 0..inter {
16336                g[i] *= u[i];
16337            }
16338        } else if d.act == Act::Silu && {
16339            let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16340            QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
16341        } {
16342            // g now holds silu(gate)·up directly.
16343        } else {
16344            u.resize(inter, 0.0);
16345            // Multi-matrix job: gate+up under one pool dispatch.
16346            let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16347            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16348            for i in 0..inter {
16349                g[i] = d.act.combine(g[i], u[i]);
16350            }
16351        }
16352        // DTG-MA bake probe (Patent 2): accumulate this layer's
16353        // per-neuron activation mass while a probe pass is active.
16354        // `CMF_FFN_PROBE_TOPK=k` switches the statistic from mass to a
16355        // HIT COUNT — how many tokens rank the neuron in their own top
16356        // k. Mass asks "how loud is this neuron overall", the count
16357        // asks "how often does this task actually need it", and the two
16358        // rank neurons differently whenever a few tokens are loud.
16359        FFN_PROBE.with(|pr| {
16360            if let Some(acc) = pr.borrow_mut().as_mut() {
16361                let li = crate::gpu::cur_layer();
16362                if li >= 0 {
16363                    if let Some(row) = acc.get_mut(li as usize) {
16364                        match probe_topk() {
16365                            0 if probe_sq() => {
16366                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16367                                    *a += (v as f64) * (v as f64);
16368                                }
16369                            }
16370                            0 if probe_signed() => {
16371                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16372                                    *a += v as f64;
16373                                }
16374                            }
16375                            0 => {
16376                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16377                                    *a += (v as f64).abs();
16378                                }
16379                            }
16380                            k => {
16381                                let n = g.len();
16382                                let k = k.min(n);
16383                                let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16384                                let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16385                                    b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16386                                });
16387                                let thr = *kth;
16388                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16389                                    if v.abs() >= thr {
16390                                        *a += 1.0;
16391                                    }
16392                                }
16393                            }
16394                        }
16395                    }
16396                }
16397            }
16398        });
16399        if oracle_topk() > 0 {
16400            keep_top_k(g, oracle_topk());
16401        }
16402        {
16403            let li = crate::gpu::cur_layer();
16404            if li >= 0 {
16405                adump_row(li as usize, g);
16406            }
16407        }
16408        let mut out = attention::take_buf(d.down_proj.rows());
16409        let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnDown);
16410        d.down_proj.matvec(g, &mut out, pool);
16411        out
16412    })
16413}
16414
16415/// Online accumulators for the AWNP refit of a narrowed FFN.
16416///
16417/// The refit needs `Gss = A_SᵀA_S` and `YA = YᵀA_S` per layer, where `A_S`
16418/// are the calibration activations of the KEPT neurons and `Y` the full
16419/// FFN output. Both are small enough to hold; the thing that is not is
16420/// the activations they are built from — a 27B layer would dump a
16421/// gigabyte per thousand tokens. So they are accumulated as the
16422/// calibration runs and written once at the end.
16423///
16424/// `CMF_FFN_REFIT=<dir>` holds `support.<L>.u32` (a u32 count then the
16425/// kept indices) for every layer to accumulate; `CMF_FFN_REFIT_FROM/TO`
16426/// bound the layer span so the accumulators fit in RAM.
16427pub struct RefitAcc {
16428    pub support: Vec<u32>,
16429    pub gss: Vec<f32>,
16430    pub ya: Vec<f32>,
16431    pub hidden: usize,
16432    pub tokens: u64,
16433    /// Activations staged transposed ([ns, t] and [hidden, t]) until the
16434    /// batch is worth a GEMM. The product costs `ns²` to move and add
16435    /// REGARDLESS of how many tokens went into it, so folding 16 chunks
16436    /// into one call cuts that cost 16× — it was 15 TB of traffic per
16437    /// calibration pass at one call per 256 tokens.
16438    pub buf_g: Vec<f32>,
16439    pub buf_o: Vec<f32>,
16440    pub buf_t: usize,
16441}
16442
16443/// The product buffer is SHARED across layers — one 473 MB allocation,
16444/// not one per layer (that was 30 GB of nothing on a 64-layer model).
16445/// It lives under the same lock as the accumulators.
16446type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
16447
16448static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
16449    std::sync::OnceLock::new();
16450
16451/// Is an FFN probe accumulator installed on this thread? The fused GPU
16452/// FFN must decline while one is, or the probe silently measures zero.
16453fn ffn_probe_active() -> bool {
16454    FFN_PROBE.with(|p| p.borrow().is_some())
16455}
16456
16457fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
16458    REFIT
16459        .get_or_init(|| {
16460            std::env::var("CMF_FFN_REFIT").ok().map(|d| {
16461                (
16462                    d,
16463                    std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
16464                )
16465            })
16466        })
16467        .as_ref()
16468}
16469
16470/// Accumulate one prefill panel into the layer's refit statistics.
16471fn refit_accumulate(
16472    li: usize,
16473    g: &[f32],
16474    b: usize,
16475    inter: usize,
16476    out: &[f32],
16477    hidden: usize,
16478    pool: Option<&Pool>,
16479) {
16480    let Some((dir, map)) = refit_dir() else {
16481        return;
16482    };
16483    static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16484    let (from, to) = *SPAN.get_or_init(|| {
16485        let g = |k: &str, d: usize| {
16486            std::env::var(k)
16487                .ok()
16488                .and_then(|v| v.parse().ok())
16489                .unwrap_or(d)
16490        };
16491        (
16492            g("CMF_FFN_REFIT_FROM", 0),
16493            g("CMF_FFN_REFIT_TO", usize::MAX),
16494        )
16495    });
16496    if li < from || li > to {
16497        return;
16498    }
16499    let mut guard = map.lock().unwrap();
16500    let (map, shared) = &mut *guard;
16501    let acc = match map.entry(li) {
16502        std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
16503        std::collections::hash_map::Entry::Vacant(e) => {
16504            let path = format!("{dir}/support.{li}.u32");
16505            let Ok(bytes) = std::fs::read(&path) else {
16506                eprintln!("refit: no {path} — layer {li} skipped");
16507                return;
16508            };
16509            let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
16510            let support: Vec<u32> = bytes[4..4 + n * 4]
16511                .chunks_exact(4)
16512                .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16513                .collect();
16514            eprintln!(
16515                "refit: layer {li} support {n} ({:.0} MB of accumulator)",
16516                (n * n + hidden * n) as f64 * 4.0 / 1e6
16517            );
16518            e.insert(RefitAcc {
16519                gss: vec![0.0; n * n],
16520                ya: vec![0.0; hidden * n],
16521                buf_g: Vec::new(),
16522                buf_o: Vec::new(),
16523                buf_t: 0,
16524                support,
16525                hidden,
16526                tokens: 0,
16527            })
16528        }
16529    };
16530    let ns = acc.support.len();
16531    // Stage this chunk transposed; the GEMM fires once the batch is full.
16532    let cap = refit_batch();
16533    if acc.buf_g.is_empty() {
16534        acc.buf_g = vec![0.0; ns * cap];
16535        acc.buf_o = vec![0.0; hidden * cap];
16536    }
16537    let take = b.min(cap - acc.buf_t);
16538    for t in 0..take {
16539        let col = acc.buf_t + t;
16540        for (j, &n) in acc.support.iter().enumerate() {
16541            acc.buf_g[j * cap + col] = g[t * inter + n as usize];
16542        }
16543        for h in 0..hidden {
16544            acc.buf_o[h * cap + col] = out[t * hidden + h];
16545        }
16546    }
16547    acc.buf_t += take;
16548    acc.tokens += take as u64;
16549    if acc.buf_t < cap {
16550        return;
16551    }
16552    let bt = acc.buf_t;
16553    acc.buf_t = 0;
16554    // The GEMM WRITES its C (it zeroes the accumulators it uses), so the
16555    // chunk product lands in scratch and is added on — the one thing that
16556    // silently turns a Gram over 13 000 tokens into a Gram over 256.
16557    // Both products are `C[n, m] += X[n, b] · Yᵀ[b, m]` with X and Y
16558    // stored row-major [·, b] — exactly `gemm_nt_f32`'s shape, so the
16559    // card does them when it is up (this is the whole calibration's
16560    // cost: O(|S|²) per token, 2.9 PFLOP for a 27B pass). The tiled CPU
16561    // loop stays as the fallback. Neither accumulates, so the product
16562    // lands in scratch and is added on.
16563    let RefitAcc {
16564        gss,
16565        ya,
16566        buf_g,
16567        buf_o,
16568        ..
16569    } = acc;
16570    let need = (ns * ns).max(hidden * ns);
16571    if shared.len() < need {
16572        shared.resize(need, 0.0);
16573    }
16574    let scratch = &mut shared[..];
16575    let _ = bt;
16576    if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
16577        add_into(gss, &scratch[..ns * ns], pool);
16578        if crate::gpu::gemm_nt_f32_transient(
16579            buf_o,
16580            buf_g,
16581            &mut scratch[..hidden * ns],
16582            hidden,
16583            cap,
16584            ns,
16585        ) {
16586            add_into(ya, &scratch[..hidden * ns], pool);
16587        } else {
16588            accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16589        }
16590    } else {
16591        accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
16592        accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
16593    }
16594    // No zeroing: the batch is always filled exactly (cap is a multiple
16595    // of the prefill chunk), and a memset of 178 MB a layer would cost
16596    // more than the GEMM.
16597}
16598
16599/// `CMF_FFN_REFIT_BATCH` — tokens staged before each GEMM (default 4096).
16600fn refit_batch() -> usize {
16601    static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16602    *B.get_or_init(|| {
16603        std::env::var("CMF_FFN_REFIT_BATCH")
16604            .ok()
16605            .and_then(|v| v.parse().ok())
16606            .unwrap_or(4096)
16607    })
16608}
16609
16610/// `c[m, n] += Σ_t left[m, t]·right[n, t]` — both operands transposed,
16611/// the CPU fallback for the staged batch.
16612fn accum_outer_t(
16613    c: &mut [f32],
16614    m: usize,
16615    n: usize,
16616    b: usize,
16617    left: &[f32],
16618    right: &[f32],
16619    pool: Option<&Pool>,
16620) {
16621    let ptr = SendMut(c.as_mut_ptr());
16622    let body = |i: usize| {
16623        let ptr = &ptr;
16624        let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
16625        for t in 0..b {
16626            let a = left[i * b + t];
16627            if a == 0.0 {
16628                continue;
16629            }
16630            for (j, o) in row.iter_mut().enumerate() {
16631                *o += a * right[j * b + t];
16632            }
16633        }
16634    };
16635    match pool {
16636        Some(p) if m > 1 => p.run_rows(m, &|s, e| {
16637            for i in s..e {
16638                body(i);
16639            }
16640        }),
16641        _ => {
16642            for i in 0..m {
16643                body(i);
16644            }
16645        }
16646    }
16647}
16648
16649/// `dst += src`, spread over the pool — at 118 M floats a layer this is
16650/// not a loop to leave on one core.
16651fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
16652    let n = dst.len().min(src.len());
16653    match pool {
16654        Some(p) if n >= 1 << 16 => {
16655            let ptr = SendMut(dst.as_mut_ptr());
16656            let f = |s: usize, e: usize| {
16657                let ptr = &ptr;
16658                for blk in s..e {
16659                    let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
16660                    for i in a..b {
16661                        unsafe { *ptr.0.add(i) += src[i] };
16662                    }
16663                }
16664            };
16665            p.run_rows(n.div_ceil(4096), &f);
16666        }
16667        _ => {
16668            for (d, v) in dst.iter_mut().zip(&src[..n]) {
16669                *d += *v;
16670            }
16671        }
16672    }
16673}
16674
16675/// `c[m, n] += Σ_t left[t, m]·right[t, n]`, with `left` stored [m, t] and
16676/// `right` [t, n]. Tiled over the rows of `c` so a tile stays in cache
16677/// while each token's `right` row streams past it once, and parallel
16678/// over tiles.
16679fn accum_outer(
16680    c: &mut [f32],
16681    m: usize,
16682    n: usize,
16683    b: usize,
16684    left: &[f32],
16685    right: &[f32],
16686    pool: Option<&Pool>,
16687) {
16688    const TILE: usize = 32;
16689    let tiles = m.div_ceil(TILE);
16690    let cp = SendMut(c.as_mut_ptr());
16691    let body = |ti: usize| {
16692        let cp = &cp;
16693        let i0 = ti * TILE;
16694        let i1 = (i0 + TILE).min(m);
16695        for t in 0..b {
16696            let r = &right[t * n..t * n + n];
16697            for i in i0..i1 {
16698                let a = left[i * b + t];
16699                if a == 0.0 {
16700                    continue;
16701                }
16702                // SAFETY: tiles partition c's rows; workers never overlap.
16703                let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
16704                for (o, v) in row.iter_mut().zip(r) {
16705                    *o += a * *v;
16706                }
16707            }
16708        }
16709    };
16710    match pool {
16711        Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
16712            for ti in s..e {
16713                body(ti);
16714            }
16715        }),
16716        _ => {
16717            for ti in 0..tiles {
16718                body(ti);
16719            }
16720        }
16721    }
16722}
16723
16724/// Write what the calibration accumulated: `gss.<L>.f32` and `ya.<L>.f32`.
16725pub fn refit_flush() -> usize {
16726    let Some((dir, map)) = refit_dir() else {
16727        return 0;
16728    };
16729    let guard = map.lock().unwrap();
16730    let mut n = 0;
16731    for (li, acc) in guard.0.iter() {
16732        // A silently truncated write here is a Gram that reshapes to
16733        // nothing an hour later — say it out loud instead.
16734        let w = |name: &str, v: &[f32]| {
16735            let path = format!("{dir}/{name}.{li}.f32");
16736            let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
16737            match std::fs::write(&path, &bytes) {
16738                Ok(()) => {}
16739                Err(e) => eprintln!(
16740                    "refit: FAILED to write {path} ({} MB): {e}",
16741                    bytes.len() / 1_000_000
16742                ),
16743            }
16744        };
16745        w("gss", &acc.gss);
16746        w("ya", &acc.ya);
16747        println!(
16748            "refit L{li}: {} support, {} tokens, hidden {}",
16749            acc.support.len(),
16750            acc.tokens,
16751            acc.hidden
16752        );
16753        n += 1;
16754    }
16755    n
16756}
16757
16758/// `CMF_FFN_ADUMP=<prefix>` — append every probed token's FFN activation
16759/// row to `<prefix>.<layer>.f16`. The co-activation record: which
16760/// neurons fire together, which is what a tube has to group if a token
16761/// is ever going to open one tube instead of sixteen.
16762fn adump_row(li: usize, g: &[f32]) {
16763    use std::io::Write as _;
16764    static FILES: std::sync::OnceLock<
16765        Option<(
16766            String,
16767            std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
16768        )>,
16769    > = std::sync::OnceLock::new();
16770    let Some((prefix, map)) = FILES
16771        .get_or_init(|| {
16772            std::env::var("CMF_FFN_ADUMP")
16773                .ok()
16774                .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
16775        })
16776        .as_ref()
16777    else {
16778        return;
16779    };
16780    // `CMF_FFN_ADUMP_FROM/_TO` narrow the dump to a layer span, so a big
16781    // calibration run fits on disk in a few passes instead of one.
16782    static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16783    let (from, to) = *SPAN.get_or_init(|| {
16784        let g = |k: &str, d: usize| {
16785            std::env::var(k)
16786                .ok()
16787                .and_then(|v| v.parse().ok())
16788                .unwrap_or(d)
16789        };
16790        (
16791            g("CMF_FFN_ADUMP_FROM", 0),
16792            g("CMF_FFN_ADUMP_TO", usize::MAX),
16793        )
16794    });
16795    if li < from || li > to {
16796        return;
16797    }
16798    let mut map = map.lock().unwrap();
16799    let f = map.entry(li).or_insert_with(|| {
16800        std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
16801    });
16802    let mut bytes = Vec::with_capacity(g.len() * 2);
16803    for v in g {
16804        bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
16805    }
16806    let _ = f.write_all(&bytes);
16807}
16808
16809/// `CMF_FFN_ORACLE_TOPK` — keep only the k largest |silu(g)·u| of each
16810/// token and zero the rest. Not a serving mode: it is the CEILING of
16811/// contextual sparsity — what a per-token router would be chasing —
16812/// measured by cheating, since the selection reads the very activations
16813/// it would have to predict.
16814fn oracle_topk() -> usize {
16815    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16816    *K.get_or_init(|| {
16817        std::env::var("CMF_FFN_ORACLE_TOPK")
16818            .ok()
16819            .and_then(|v| v.parse().ok())
16820            .unwrap_or(0)
16821    })
16822}
16823
16824/// `CMF_FFN_GATE_TOPK` — the REALIZABLE cousin of the oracle: rank the
16825/// neurons by their gate alone (which the kernel has computed anyway
16826/// before it reads `up`), keep the k best, and drop the rest. Every
16827/// dropped neuron's `up` row and `down` column stay unread, so this is
16828/// the sparsity a serving path can actually take without a router.
16829fn gate_topk() -> usize {
16830    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16831    *K.get_or_init(|| {
16832        std::env::var("CMF_FFN_GATE_TOPK")
16833            .ok()
16834            .and_then(|v| v.parse().ok())
16835            .unwrap_or(0)
16836    })
16837}
16838
16839/// `CMF_FFN_GATE_BLOCK` — select in blocks of B neurons instead of one
16840/// by one. A scattered per-neuron choice cannot be read efficiently (a
16841/// row at a time, no prefetch runway); a block of 32 is a contiguous
16842/// 32-row slab of `up` and of the transposed `down`, which the ordinary
16843/// kernels stream. The question the measurement answers is what the
16844/// block costs in quality.
16845fn gate_block() -> usize {
16846    static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16847    *B.get_or_init(|| {
16848        std::env::var("CMF_FFN_GATE_BLOCK")
16849            .ok()
16850            .and_then(|v| v.parse().ok())
16851            .unwrap_or(1)
16852    })
16853}
16854
16855/// Zero all but the `k` largest BLOCKS (by summed square) of a row.
16856fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
16857    let n = g.len();
16858    let nb = n.div_ceil(block);
16859    let kb = (keep_n.div_ceil(block)).clamp(1, nb);
16860    if kb >= nb {
16861        return;
16862    }
16863    let mut score: Vec<f32> = (0..nb)
16864        .map(|b| {
16865            g[b * block..((b + 1) * block).min(n)]
16866                .iter()
16867                .map(|v| v * v)
16868                .sum::<f32>()
16869        })
16870        .collect();
16871    let mut ord = score.clone();
16872    let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
16873        b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16874    });
16875    let thr = *kth;
16876    for b in 0..nb {
16877        if score[b] < thr {
16878            g[b * block..((b + 1) * block).min(n)].fill(0.0);
16879        }
16880    }
16881    score.clear();
16882}
16883
16884/// Zero all but the `k` largest magnitudes of one token's activation row.
16885fn keep_top_k(g: &mut [f32], k: usize) {
16886    if gate_block() > 1 {
16887        return keep_top_blocks(g, k, gate_block());
16888    }
16889    let n = g.len();
16890    if k == 0 || k >= n {
16891        return;
16892    }
16893    let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16894    let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16895        b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16896    });
16897    let thr = *kth;
16898    for v in g.iter_mut() {
16899        if v.abs() < thr {
16900            *v = 0.0;
16901        }
16902    }
16903}
16904
16905/// `CMF_FFN_PROBE_SQ` — accumulate Σa², so the dump divided by the token
16906/// count and square-rooted is the RMS activation trace Patent 12 weights
16907/// its matrices by.
16908fn probe_sq() -> bool {
16909    static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16910    *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
16911}
16912
16913/// `CMF_FFN_PROBE_SIGNED` — accumulate the SIGNED activation sum
16914/// instead of its magnitude: what a dropped neuron contributes ON
16915/// AVERAGE, which is the bias a narrowed FFN can add back for free.
16916fn probe_signed() -> bool {
16917    static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16918    *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
16919}
16920
16921/// `CMF_FFN_MEANFILL=<file>` — a masked-out neuron contributes its MEAN
16922/// activation instead of zero (`u32 layers, u32 inter, f32[…]`, the mass
16923/// dump layout, holding per-neuron means). Dropping a neuron outright
16924/// also drops its average contribution, which shifts the layer output by
16925/// a constant; filling the mean back is one add per layer and costs no
16926/// bytes off the bus. This is the measurement arm — in a tube file the
16927/// same correction ships as a per-task bias vector.
16928fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
16929    static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
16930    M.get_or_init(|| {
16931        let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
16932        let b = std::fs::read(&p).ok()?;
16933        let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
16934        let vals: Vec<f32> = b[8..]
16935            .chunks_exact(4)
16936            .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16937            .collect();
16938        eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
16939        Some((inter, vals))
16940    })
16941    .as_ref()
16942}
16943
16944/// `CMF_FFN_PROBE_TOPK` — 0 (default) = accumulate mass, k>0 = count
16945/// how often a neuron lands in a token's top k.
16946fn probe_topk() -> usize {
16947    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16948    *K.get_or_init(|| {
16949        std::env::var("CMF_FFN_PROBE_TOPK")
16950            .ok()
16951            .and_then(|v| v.parse().ok())
16952            .unwrap_or(0)
16953    })
16954}
16955
16956thread_local! {
16957    /// DTG-MA activation probe: per-layer per-neuron Σ|silu(g)·u|
16958    /// accumulator, alive only during `Pipeline::probe_ffn_mass`.
16959    static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
16960        const { std::cell::RefCell::new(None) };
16961}
16962
16963/// Per-token structured sparsity, paid for in bytes.
16964///
16965/// The gate is the cheapest third of an FFN and it already says which
16966/// neurons matter: `silu(gate)` near zero means the neuron contributes
16967/// nothing whatever `up` says. So compute every gate, keep the `k`
16968/// loudest, and read ONLY those neurons' `up` rows and `down` rows —
16969/// the latter needs `down_proj` stored transposed, otherwise a neuron's
16970/// down weights are a strided column and "reading only those" costs a
16971/// full cache line each.
16972///
16973/// Returns `None` when the file has no transposed `down` (the caller
16974/// then runs the ordinary dense path).
16975fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
16976    // The scatter path reads individual rows/columns and cannot express the
16977    // per-matrix signed FWHT boundary.  Let the descriptor-aware dense path
16978    // handle Prism files rather than silently running an unrotated sparse
16979    // approximation.
16980    if d.gate_proj.has_prism_contract()
16981        || d.up_proj.has_prism_contract()
16982        || d.down_proj.has_prism_contract()
16983    {
16984        return None;
16985    }
16986    let dt = d.down_t.as_ref()?;
16987    let inter = d.gate_proj.rows();
16988    let hidden = dt.cols();
16989    if k == 0 || k >= inter || d.act != Act::Silu {
16990        return None;
16991    }
16992    DYN_SCRATCH.with(|sc| {
16993        let mut sc = sc.borrow_mut();
16994        let DynScratch {
16995            g,
16996            mag,
16997            live,
16998            parts,
16999        } = &mut *sc;
17000        g.resize(inter, 0.0);
17001        d.gate_proj.matvec(x, g, pool);
17002        for v in g.iter_mut() {
17003            *v = inference::silu(*v);
17004        }
17005        // The k-th largest |silu(gate)| is the threshold; ties keep more,
17006        // which is the safe side.
17007        mag.clear();
17008        mag.extend(g.iter().map(|v| v.abs()));
17009        let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
17010            b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17011        });
17012        let thr = *kth;
17013        live.clear();
17014        live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
17015        let mut out = vec![0.0f32; hidden];
17016        match pool {
17017            Some(p) if live.len() >= 64 => {
17018                let nw = p.n_workers() + 1;
17019                parts.clear();
17020                parts.resize(nw * hidden, 0.0);
17021                let ptr = SendMut(parts.as_mut_ptr());
17022                let n = live.len();
17023                let live_ref: &[u32] = live;
17024                let g_ref: &[f32] = g;
17025                p.run(&|w, workers| {
17026                    let chunk = n.div_ceil(workers);
17027                    let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
17028                    if s >= e {
17029                        return;
17030                    }
17031                    WORKER_SCRATCH.with(|ws| {
17032                        let mut ws = ws.borrow_mut();
17033                        let [scratch, acc] = &mut *ws;
17034                        scratch.resize(hidden.max(x.len()), 0.0);
17035                        acc.clear();
17036                        acc.resize(hidden, 0.0);
17037                        for (o, &nrm) in live_ref[s..e].iter().enumerate() {
17038                            // One neuron of runway: the next row's lines
17039                            // start moving while this one is multiplied.
17040                            if let Some(&nx) = live_ref[s..e].get(o + 1) {
17041                                d.up_proj.prefetch_row(nx as usize);
17042                                dt.prefetch_row(nx as usize);
17043                            }
17044                            let idx = nrm as usize;
17045                            let up = d.up_proj.row_dot(idx, x, scratch);
17046                            let a = g_ref[idx] * up;
17047                            if a != 0.0 {
17048                                dt.add_row_scaled(idx, a, acc, scratch);
17049                            }
17050                        }
17051                        for (j, v) in acc.iter().enumerate() {
17052                            unsafe { *ptr.at(w * hidden + j) = *v };
17053                        }
17054                    });
17055                });
17056                for w in 0..nw {
17057                    for (j, o) in out.iter_mut().enumerate() {
17058                        *o += parts[w * hidden + j];
17059                    }
17060                }
17061            }
17062            _ => {
17063                WORKER_SCRATCH.with(|ws| {
17064                    let mut ws = ws.borrow_mut();
17065                    let [scratch, _acc] = &mut *ws;
17066                    scratch.resize(hidden.max(x.len()), 0.0);
17067                    for &nrm in live.iter() {
17068                        let idx = nrm as usize;
17069                        let up = d.up_proj.row_dot(idx, x, scratch);
17070                        let a = g[idx] * up;
17071                        if a != 0.0 {
17072                            dt.add_row_scaled(idx, a, &mut out, scratch);
17073                        }
17074                    }
17075                });
17076            }
17077        }
17078        Some(out)
17079    })
17080}
17081
17082/// Caller-side scratch of the dynamic path — one allocation per thread,
17083/// not one per layer per token (that alone cost a third of the decode).
17084struct DynScratch {
17085    g: Vec<f32>,
17086    mag: Vec<f32>,
17087    live: Vec<u32>,
17088    parts: Vec<f32>,
17089}
17090
17091thread_local! {
17092    static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
17093        std::cell::RefCell::new(DynScratch {
17094            g: Vec::new(),
17095            mag: Vec::new(),
17096            live: Vec::new(),
17097            parts: Vec::new(),
17098        })
17099    };
17100    /// Pool-worker scratch: the row buffer and this worker's partial sum.
17101    static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
17102        const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
17103}
17104
17105/// `dense_ffn_cpu` with a per-visit mask landing on the activations —
17106/// the masked-inference fast path's decode arm. Full fused quant
17107/// compute, closed neurons zeroed before down: arithmetically the
17108/// pruned network, no dequant, no weight bytes touched.
17109fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
17110    let inter = d.gate_proj.rows();
17111    FFN_SCRATCH.with(|s| {
17112        let mut s = s.borrow_mut();
17113        let [g, u, ..] = &mut *s;
17114        g.resize(inter, 0.0);
17115        if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
17116            // g holds silu(gate)·up.
17117        } else {
17118            u.resize(inter, 0.0);
17119            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
17120            for i in 0..inter {
17121                g[i] = d.act.combine(g[i], u[i]);
17122            }
17123        }
17124        zero_masked_cols(g, 1, inter, mask_row);
17125        let mut out = attention::take_buf(d.down_proj.rows());
17126        d.down_proj.matvec(g, &mut out, pool);
17127        out
17128    })
17129}
17130
17131/// Dense FFN as one GPU submission via the MoE block path (single
17132/// expert, weight 1.0): gate → silu·up → down chained in one command
17133/// buffer, intermediate activations device-resident. None → weights
17134/// not q8-mapped in the primary shard / over the VRAM budget / backend
17135/// refusal → honest CPU path.
17136fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
17137    if d.gate_proj.has_prism_contract()
17138        || d.up_proj.has_prism_contract()
17139        || d.down_proj.has_prism_contract()
17140    {
17141        return None;
17142    }
17143    // The GPU block hardcodes SiLU; GeLU FFNs (Gemma) stay on CPU.
17144    if d.act != Act::Silu {
17145        return None;
17146    }
17147    // Threshold: tiny FFNs are not worth a submission (q1 excepted —
17148    // see the caller's gate).
17149    if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
17150        return None;
17151    }
17152    let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
17153    let mut model_ref = None;
17154    moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
17155    let model = model_ref?;
17156    let hidden = jobs[0].down.1;
17157    let mut out = attention::take_buf(hidden);
17158    if crate::gpu::moe_block(&model, &jobs, &mut out) {
17159        Some(out)
17160    } else {
17161        let mut out = out;
17162        attention::recycle_buf(&mut out);
17163        None
17164    }
17165}
17166
17167/// q8-mapped primary-shard tensor parts for a GPU job: q8_2f carries
17168/// its column field, q8_row runs with empty col slices (the backend
17169/// skips the multiply). Shared by the MoE block and the dense-FFN
17170/// single-job path.
17171#[allow(clippy::type_complexity)]
17172#[allow(clippy::type_complexity)]
17173pub(crate) fn moe_parts(
17174    t: &QTensor,
17175) -> Option<(
17176    &std::sync::Arc<cortiq_core::CmfModel>,
17177    usize,
17178    usize,
17179    usize,
17180    &[f32],
17181    &[f32],
17182    bool,
17183    bool,
17184    bool,
17185)> {
17186    match t {
17187        QTensor::Mapped {
17188            model,
17189            idx,
17190            dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
17191            rows,
17192            cols,
17193            row_scale,
17194            col_field,
17195            ..
17196        } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
17197            model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
17198        )),
17199        // q1: tile-embedded scales — empty rs/col slices, raw xs.
17200        QTensor::Mapped {
17201            model,
17202            idx,
17203            dtype: cortiq_core::TensorDtype::Q1,
17204            rows,
17205            cols,
17206            ..
17207        } => Some((
17208            model,
17209            *idx,
17210            *rows,
17211            *cols,
17212            &[][..],
17213            &[][..],
17214            true,
17215            false,
17216            false,
17217        )),
17218        // q4_tiled: 18-byte tiles with embedded f16 scales — raw xs.
17219        QTensor::Mapped {
17220            model,
17221            idx,
17222            dtype: cortiq_core::TensorDtype::Q4Tiled,
17223            rows,
17224            cols,
17225            ..
17226        } => Some((
17227            model,
17228            *idx,
17229            *rows,
17230            *cols,
17231            &[][..],
17232            &[][..],
17233            false,
17234            true,
17235            false,
17236        )),
17237        // q4tp: same raw-xs contract, different stride and scale plane.
17238        QTensor::Mapped {
17239            model,
17240            idx,
17241            dtype: cortiq_core::TensorDtype::Q4TiledP,
17242            rows,
17243            cols,
17244            ..
17245        } => Some((
17246            model,
17247            *idx,
17248            *rows,
17249            *cols,
17250            &[][..],
17251            &[][..],
17252            false,
17253            true,
17254            false,
17255        )),
17256        // q2tp: the 2-bit expert plane of the mixed profile — q4 family
17257        // for stride bookkeeping, flagged q2 so the trio validation can
17258        // demand a q4tp down.
17259        QTensor::Mapped {
17260            model,
17261            idx,
17262            dtype: cortiq_core::TensorDtype::Q2TiledP,
17263            rows,
17264            cols,
17265            ..
17266        } => Some((
17267            model,
17268            *idx,
17269            *rows,
17270            *cols,
17271            &[][..],
17272            &[][..],
17273            false,
17274            true,
17275            true,
17276        )),
17277        _ => None,
17278    }
17279}
17280
17281/// Map a MoE onto the Metal token graph's contract: f32 router, a
17282/// shared expert (gated — Qwen — or ungated at weight 1 — DeepSeek-V3 /
17283/// HunYuan hy_v3), softmax or sigmoid scores with an optional selection
17284/// bias and routed scale, experts uniformly q4tp (or the mixed profile:
17285/// q2tp gate/up over a q4tp down). τ routers, masks, per-expert scales
17286/// and Gemma's router-input norm refuse here — those semantics stay on
17287/// the CPU path.
17288#[cfg(target_os = "macos")]
17289fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
17290    if m.router_input_norm
17291        || m.route_tau.is_some()
17292        || m.mask.is_some()
17293        || m.per_expert_scale.is_some()
17294        || m.experts.is_empty()
17295        || m.top_k == 0
17296        || m.resonance.is_some()
17297    {
17298        return None;
17299    }
17300    // The select kernel always fills the shared slot: a model without a
17301    // shared expert (LFM2-MoE) stays on the CPU path here.
17302    let (sh, sg) = match &m.shared {
17303        Some((sh, sg)) => (sh, sg.as_ref()),
17304        None => return None,
17305    };
17306    let (rf, rr, rc) = m.router.f32_parts()?;
17307    if rr != m.experts.len() || rc != hidden {
17308        return None;
17309    }
17310    let shared_gated = sg.is_some();
17311    let sf = match sg {
17312        Some(sg) => {
17313            let (sf, sr, sc) = sg.f32_parts()?;
17314            if sr * sc != hidden {
17315                return None;
17316            }
17317            sf
17318        }
17319        // Ungated: the router's first row stands in for the gate matvec
17320        // (its logit is never read — the kernel pins weight 1).
17321        None => &rf[..hidden],
17322    };
17323    if let Some(b) = &m.expert_bias {
17324        if b.len() != m.experts.len() {
17325            return None;
17326        }
17327    }
17328    let inter = m.experts[0].gate_proj.rows();
17329    // The first expert's gate decides the profile; every trio (shared
17330    // included) must agree — the jobs ladder flips ONE kernel for all.
17331    let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
17332    let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
17333        if e.act != Act::Silu
17334            || e.gate_proj.rows() != inter
17335            || e.gate_proj.cols() != hidden
17336            || e.up_proj.rows() != inter
17337            || e.up_proj.cols() != hidden
17338            || e.down_proj.rows() != hidden
17339            || e.down_proj.cols() != inter
17340        {
17341            return None;
17342        }
17343        let pick = |t: &QTensor| -> Option<usize> {
17344            if gu_q2 {
17345                t.mapped_q2tp().map(|(_, i)| i)
17346            } else {
17347                t.mapped_q4tp().map(|(_, i)| i)
17348            }
17349        };
17350        Some((
17351            pick(&e.gate_proj)?,
17352            pick(&e.up_proj)?,
17353            e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
17354        ))
17355    };
17356    let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
17357    let shared = trio(sh)?;
17358    Some(crate::gpu::GpuMoe {
17359        router: rf,
17360        sgate: sf,
17361        experts,
17362        shared,
17363        n_exp: m.experts.len(),
17364        top_k: m.top_k,
17365        inter,
17366        norm_topk: m.norm_topk_prob,
17367        route_scale: m.routed_scaling,
17368        gu_q2,
17369        sigmoid: m.router_sigmoid,
17370        bias: m.expert_bias.as_deref(),
17371        shared_gated,
17372    })
17373}
17374
17375/// Build one gate/up/down GPU job from three tensors. `moe_push_job` is the
17376/// DenseFfn-shaped caller; architectures that keep their experts in their own
17377/// structs (DeepSeek-V4) come here directly.
17378pub(crate) fn moe_push_job_parts<'a>(
17379    gate: &'a QTensor,
17380    up: &'a QTensor,
17381    down: &'a QTensor,
17382    x: &[f32],
17383    w: f32,
17384    swiglu_limit: 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    let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
17390    let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
17391    let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
17392    if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17393        return None; // mixed-dtype trio — honest CPU path
17394    }
17395    // The 2-bit profile is gate/up q2tp over a PLAIN q4tp down; any other
17396    // 2-bit arrangement stays on the CPU.
17397    if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
17398        return None;
17399    }
17400    if !gq2 && dq2 {
17401        return None;
17402    }
17403    model_ref.get_or_insert_with(|| gm.clone());
17404    let dt = |cf: &[f32]| {
17405        if cf.is_empty() {
17406            cortiq_core::TensorDtype::Q8Row
17407        } else {
17408            cortiq_core::TensorDtype::Q8_2f
17409        }
17410    };
17411    jobs.push(crate::gpu::MoeJob {
17412        gate: (gi, gr, gc, grs),
17413        up: (ui, ur, uc, urs),
17414        down: (di, dr, dc, drs),
17415        xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
17416        xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
17417        down_col: dcf,
17418        w,
17419        q1: gq1,
17420        q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
17421        q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
17422        gu_q2: gq2,
17423        swiglu_limit,
17424    });
17425    Some(())
17426}
17427
17428/// Build one gate/up/down GPU job (see `moe_parts`).
17429fn moe_push_job<'a>(
17430    d: &'a DenseFfn,
17431    x: &[f32],
17432    w: f32,
17433    jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17434    model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17435) -> Option<()> {
17436    use crate::qtensor::prescale;
17437    if d.act != Act::Silu {
17438        return None; // GPU block hardcodes SiLU
17439    }
17440    let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
17441    let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
17442    let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
17443    if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17444        return None; // mixed-dtype trio — honest CPU path
17445    }
17446    if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
17447        return None;
17448    }
17449    if !gq2 && dq2 {
17450        return None;
17451    }
17452    model_ref.get_or_insert_with(|| gm.clone());
17453    let gdt = if gcf.is_empty() {
17454        cortiq_core::TensorDtype::Q8Row
17455    } else {
17456        cortiq_core::TensorDtype::Q8_2f
17457    };
17458    let udt = if ucf.is_empty() {
17459        cortiq_core::TensorDtype::Q8Row
17460    } else {
17461        cortiq_core::TensorDtype::Q8_2f
17462    };
17463    jobs.push(crate::gpu::MoeJob {
17464        gate: (gi, gr, gc, grs),
17465        up: (ui, ur, uc, urs),
17466        down: (di, dr, dc, drs),
17467        xs_gate: prescale(x, gcf, gdt).into_owned(),
17468        xs_up: prescale(x, ucf, udt).into_owned(),
17469        down_col: dcf,
17470        w,
17471        q1: gq1,
17472        q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
17473        q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
17474        gu_q2: gq2,
17475        swiglu_limit: 0.0,
17476    });
17477    Some(())
17478}
17479
17480/// Sparse dense-FFN directly on QUANTIZED weights (mask × mmap): reads
17481/// ONLY the active neurons' gate/up rows and down columns from the mmap
17482/// — no full-matrix dequant, no f32 model copy. This is what lets a
17483/// masked big model run at quantized RSS (the historical mask path
17484/// forced the whole model to f32). Semantics identical to the f32
17485/// sparse path within quant tolerance.
17486fn sparse_ffn_quant(
17487    d: &DenseFfn,
17488    x: &[f32],
17489    active: &[u16],
17490    hidden: usize,
17491    pool: Option<&Pool>,
17492) -> Vec<f32> {
17493    let n = active.len();
17494    let inter = d.gate_proj.rows();
17495    let mut act = vec![0.0f32; n];
17496    // Scratch is needed if EITHER projection is group-packed (q4/vbit);
17497    // gate/up normally share a dtype but sizing on both is robust.
17498    let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
17499    let compute = |ai: usize| -> f32 {
17500        let idx = active[ai] as usize;
17501        if idx >= inter {
17502            return 0.0; // defensive parity with the f32 sparse path
17503        }
17504        let mut s = if need_scratch {
17505            vec![0.0f32; hidden]
17506        } else {
17507            Vec::new()
17508        };
17509        let gate = d.gate_proj.row_dot(idx, x, &mut s);
17510        let up = d.up_proj.row_dot(idx, x, &mut s);
17511        d.act.combine(gate, up)
17512    };
17513    match pool {
17514        Some(p) if n >= 256 => {
17515            let ptr = SendMut(act.as_mut_ptr());
17516            p.run(&|widx, nw| {
17517                let chunk = n.div_ceil(nw);
17518                let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
17519                for ai in s..e {
17520                    unsafe { *ptr.at(ai) = compute(ai) };
17521                }
17522            });
17523        }
17524        _ => {
17525            for (ai, a) in act.iter_mut().enumerate() {
17526                *a = compute(ai);
17527            }
17528        }
17529    }
17530    // Scatter through active down columns (reads only those columns).
17531    let mut out = vec![0.0f32; hidden];
17532    for (ai, &idx) in active.iter().enumerate() {
17533        let w = act[ai];
17534        if w.abs() >= 1e-12 && (idx as usize) < inter {
17535            d.down_proj.add_col_scaled(idx as usize, w, &mut out);
17536        }
17537    }
17538    out
17539}
17540
17541/// Test-only re-export of the private sparse-quant FFN (mask × mmap gate).
17542#[doc(hidden)]
17543pub fn sparse_ffn_quant_for_test(
17544    d: &DenseFfn,
17545    x: &[f32],
17546    active: &[u16],
17547    hidden: usize,
17548) -> Vec<f32> {
17549    sparse_ffn_quant(d, x, active, hidden, None)
17550}
17551
17552/// Dequantize a DenseFfn's three matrices to f32 (transient; only the
17553/// q4/vbit-masked fallback uses it — the memory-lean path is
17554/// sparse_ffn_quant). Reuses row_f32 row-by-row.
17555fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
17556    let deq = |t: &QTensor| -> Vec<f32> {
17557        let (rows, cols) = (t.rows(), t.cols());
17558        let mut out = vec![0.0f32; rows * cols];
17559        for r in 0..rows {
17560            t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
17561        }
17562        out
17563    };
17564    (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
17565}
17566
17567/// Pointer wrapper for the worker-pool scatter (same pattern as qtensor).
17568struct SendMut(*mut f32);
17569unsafe impl Send for SendMut {}
17570unsafe impl Sync for SendMut {}
17571impl SendMut {
17572    #[inline]
17573    // Deliberate unsynchronized scatter: pool workers write disjoint indices
17574    // in parallel, so returning `&mut` from `&self` is intentional here.
17575    #[allow(clippy::mut_from_ref)]
17576    unsafe fn at(&self, i: usize) -> &mut f32 {
17577        unsafe { &mut *self.0.add(i) }
17578    }
17579}
17580
17581/// Router → (selected experts in torch.topk order, per-expert score
17582/// vector, normalizer). The final weight of expert `e` is `p[e] / wsum`.
17583///
17584/// Two regimes share this. Qwen: softmax over ALL experts, top-k of the
17585/// probabilities, optional renorm — `router_sigmoid=false`, no bias,
17586/// scale 1 → bit-identical to the historical path. LFM2-MoE /
17587/// DeepSeek-V3 `noaux_tc`: per-expert sigmoid scores, an optional
17588/// selection bias (top-k CHOICE only; weights stay unbiased), a 1e-6 renorm
17589/// floor and a routed scale. Architectures whose reference uses a different
17590/// sigmoid denominator floor (for example GLM-5's `1e-20`) call
17591/// [`moe_route_with_eps`] directly; the historical generic path remains
17592/// unchanged.
17593pub(crate) fn moe_route(
17594    logits: &[f32],
17595    m: &MoeFfn,
17596    allowed: Option<&[bool]>,
17597) -> (Vec<usize>, Vec<f32>, f32) {
17598    moe_route_with_eps(logits, m, allowed, 1e-6)
17599}
17600
17601/// Router implementation with an explicit sigmoid renormalization floor.
17602///
17603/// GLM-5.3's source computes `sum(selected_scores) + 1e-20`; using the
17604/// generic 1e-6 floor there is not a harmless tolerance difference when all
17605/// logits are very negative: it collapses the routed branch toward zero
17606/// instead of normalizing the selected experts. Keeping the epsilon parameter
17607/// here avoids changing the established Qwen/LFM2 contract while allowing
17608/// each architecture to preserve its own numerical semantics.
17609pub(crate) fn moe_route_with_eps(
17610    logits: &[f32],
17611    m: &MoeFfn,
17612    allowed: Option<&[bool]>,
17613    sigmoid_denom_eps: f32,
17614) -> (Vec<usize>, Vec<f32>, f32) {
17615    let ne = logits.len();
17616    // Expert restriction: the static env mask (CMF_MOE_MASK) AND the
17617    // active task mask's expert fields (spec §5) both narrow the
17618    // candidate set; selection happens over the admitted experts only.
17619    // With norm_topk the kept weights renormalize below; without it
17620    // the excluded mass is honestly dropped.
17621    let admit = |e: usize| {
17622        m.mask.as_ref().is_none_or(|mk| mk[e])
17623            && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
17624    };
17625    // The resonance router (spec §9.5.1) selects by the RAW score: the
17626    // trainer (`resonance_winner`) and the resident graph
17627    // (`embryo_core_route_pick`) take the first maximum of the scores
17628    // and run the winner with weight 1.0. Selecting through the softmax
17629    // instead is not the same decision: `exp(l − max)` rounds two scores
17630    // closer than 2^-25 (possible below |score| 0.25) to the same 1.0, and
17631    // the lower index would take a token whose score is strictly smaller
17632    // — the trainer's trace and the graph would disagree with this path.
17633    // `−∞` (outside the shell) never wins; with no finite admitted expert
17634    // the generic path below degrades to uniform. The winner's
17635    // probability is 1.0 by construction (a one-hot `p`), so its
17636    // renormalized weight is `routed_scaling` on both norm_topk settings.
17637    if m.resonance.is_some() && m.top_k == 1 {
17638        let mut best: Option<usize> = None;
17639        for e in (0..ne).filter(|&e| admit(e)) {
17640            let l = logits[e];
17641            if l == f32::NEG_INFINITY || l.is_nan() {
17642                continue;
17643            }
17644            if best.is_none_or(|b| l > logits[b]) {
17645                best = Some(e);
17646            }
17647        }
17648        if let Some(b) = best {
17649            let mut p = vec![0.0f32; ne];
17650            p[b] = 1.0;
17651            return (vec![b], p, 1.0 / m.routed_scaling);
17652        }
17653    }
17654    // A `−∞` logit (a grown expert outside its shell, `Resonance::scores`)
17655    // takes probability 0 on both paths: sigmoid(−∞) = 0, exp(−∞ − max) = 0
17656    // — top-1 is the best FINITE expert, its renormalized weight exactly
17657    // 1.0. Every expert at −∞ cannot happen (trunk experts have no shell);
17658    // should it, the softmax would be NaN, so it degrades to uniform.
17659    let p: Vec<f32> = if m.router_sigmoid {
17660        logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
17661    } else {
17662        let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
17663        if mx == f32::NEG_INFINITY {
17664            vec![1.0 / ne.max(1) as f32; ne]
17665        } else {
17666            let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
17667            let s: f32 = e.iter().sum();
17668            for v in &mut e {
17669                *v /= s;
17670            }
17671            e
17672        }
17673    };
17674    let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
17675    // Descending by selection score, lower index wins ties (torch.topk).
17676    match &m.expert_bias {
17677        Some(b) => idx.sort_unstable_by(|&x, &y| {
17678            (p[y] + b[y])
17679                .partial_cmp(&(p[x] + b[x]))
17680                .unwrap()
17681                .then(x.cmp(&y))
17682        }),
17683        None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
17684    }
17685    idx.truncate(m.top_k);
17686    // Adaptive τ-routing: trim the tail experts once the kept mass is
17687    // enough. wsum below renormalizes over the KEPT set, so the output
17688    // stays a proper weighted average.
17689    if let Some(tau) = m.route_tau {
17690        let total: f32 = idx.iter().map(|&e| p[e]).sum();
17691        if total > 0.0 {
17692            let mut acc = 0.0f32;
17693            let mut keep = idx.len();
17694            for (i, &e) in idx.iter().enumerate() {
17695                acc += p[e];
17696                if acc >= tau * total {
17697                    keep = i + 1;
17698                    break;
17699                }
17700            }
17701            idx.truncate(keep);
17702        }
17703    }
17704    let wsum: f32 = if m.norm_topk_prob {
17705        let s: f32 = idx.iter().map(|&e| p[e]).sum();
17706        // Sigmoid routers use their architecture's reference floor; the
17707        // softmax path's probs already sum near 1, so it stays exactly as
17708        // before.
17709        (if m.router_sigmoid {
17710            s + sigmoid_denom_eps
17711        } else {
17712            s
17713        }) / m.routed_scaling
17714    } else {
17715        1.0 / m.routed_scaling
17716    };
17717    (idx, p, wsum)
17718}
17719
17720/// See the call site: one `layer:e1,e2,…` line per routed token.
17721fn moe_trace(idx: &[usize]) {
17722    moe_trace_at(crate::gpu::cur_layer() as i32, idx)
17723}
17724
17725/// The same, for callers that know their layer (DSV4 owns its layers and
17726/// never sets the pipeline's current-layer marker).
17727pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
17728    use std::io::Write;
17729    static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
17730        std::sync::OnceLock::new();
17731    let Some(f) = F.get_or_init(|| {
17732        let p = std::env::var("CMF_MOE_TRACE").ok()?;
17733        Some(std::sync::Mutex::new(
17734            std::fs::OpenOptions::new()
17735                .create(true)
17736                .append(true)
17737                .open(p)
17738                .ok()?,
17739        ))
17740    }) else {
17741        return;
17742    };
17743    let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
17744    let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
17745}
17746
17747/// MoE FFN: router → top-k experts (see `moe_route`). Only selected
17748/// experts' pages are touched in mmap.
17749pub(crate) fn moe_ffn(
17750    m: &MoeFfn,
17751    x: &[f32],
17752    pool: Option<&Pool>,
17753    allowed: Option<&[bool]>,
17754) -> Vec<f32> {
17755    let r = moe_ffn_route(m, x, pool, allowed);
17756    moe_ffn_experts(m, x, &r, pool)
17757}
17758
17759/// One token's host route through a MoE layer: the chosen experts in
17760/// selection order, the per-expert scores and the normalizer (see
17761/// `moe_route`), plus the raw router logits.
17762pub(crate) struct MoeRoute {
17763    pub idx: Vec<usize>,
17764    pub p: Vec<f32>,
17765    pub wsum: f32,
17766    pub logits: Vec<f32>,
17767}
17768
17769/// The routing half of `moe_ffn`, shared by every executor of the chosen
17770/// experts (the host/per-op path below and the MiMo dynamic device cache,
17771/// `crate::mimo_moe`): activation accounting, router logits, `moe_route`,
17772/// the selection statistics and the `CMF_MOE_TRACE` line — so switching
17773/// executors can never change which experts a token gets.
17774pub(crate) fn moe_ffn_route(
17775    m: &MoeFfn,
17776    x: &[f32],
17777    pool: Option<&Pool>,
17778    allowed: Option<&[bool]>,
17779) -> MoeRoute {
17780    accumulate_act(m, x, 1);
17781    let ne = m.experts.len();
17782    let mut logits = vec![0.0f32; ne];
17783    match &m.resonance {
17784        Some(r) => r.scores(x, &mut logits),
17785        None => m.router.matvec(x, &mut logits, pool),
17786    }
17787    let (idx, p, wsum) = moe_route(&logits, m, allowed);
17788    {
17789        let mut st = m.stats.borrow_mut();
17790        if st.len() < ne {
17791            st.resize(ne, 0);
17792        }
17793        for &e in &idx {
17794            st[e] += 1;
17795        }
17796    }
17797    // `CMF_MOE_TRACE=<file>`: append one line per (layer, token) with the
17798    // selected expert ids. The cumulative `stats` above answer "which
17799    // experts are popular"; a residency design needs the question they
17800    // cannot answer — whether CONSECUTIVE tokens reuse experts (the
17801    // temporal locality an LRU cache lives on, FreeToken §4).
17802    moe_trace(&idx);
17803    MoeRoute {
17804        idx,
17805        p,
17806        wsum,
17807        logits,
17808    }
17809}
17810
17811/// The expert half of `moe_ffn`: run a route's experts on the per-op GPU
17812/// block or the host.
17813pub(crate) fn moe_ffn_experts(
17814    m: &MoeFfn,
17815    x: &[f32],
17816    r: &MoeRoute,
17817    pool: Option<&Pool>,
17818) -> Vec<f32> {
17819    let (idx, p, wsum) = (&r.idx, &r.p, r.wsum);
17820    // D5: the whole layer MoE block in one GPU command buffer (experts — the
17821    // same mmap via a no-copy buffer; intermediate activations on the GPU).
17822    // Same Ffn probe class as the dense chain: one submit per layer
17823    // either wins on this driver stack or it doesn't.
17824    if crate::gpu::enabled_here() {
17825        match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
17826            crate::gpu::ProbeArm::Gpu => {
17827                let t0 = std::time::Instant::now();
17828                if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
17829                    crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
17830                    return out;
17831                }
17832            }
17833            crate::gpu::ProbeArm::CpuTimed => {
17834                let t0 = std::time::Instant::now();
17835                let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
17836                crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
17837                return out;
17838            }
17839            crate::gpu::ProbeArm::Cpu => {
17840                return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
17841            }
17842        }
17843    }
17844    moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
17845}
17846
17847/// One MoE token through the MiMo expert bank (`crate::mimo_moe`), or —
17848/// when the bank does not serve it — through the host path with the SAME
17849/// route, so the routing statistics and `CMF_MOE_TRACE` see it once.
17850fn moe_ffn_banked(
17851    slot: &mut crate::mimo_moe::Slot,
17852    li: usize,
17853    m: &MoeFfn,
17854    x: &[f32],
17855    pool: Option<&Pool>,
17856) -> Vec<f32> {
17857    let t0 = std::time::Instant::now();
17858    let r = moe_ffn_route(m, x, pool, None);
17859    slot.note_route(t0.elapsed().as_nanos() as u64);
17860    match slot.forward(li, m, x, &r, pool) {
17861        Some(out) => out,
17862        None => crate::qtensor::float_activations_scope(|| {
17863            crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, &r, pool))
17864        }),
17865    }
17866}
17867
17868/// Verify rows share a bank frame; routing and fallback are decode's.
17869fn moe_ffn_banked_rows(
17870    slot: &mut crate::mimo_moe::Slot,
17871    li: usize,
17872    m: &MoeFfn,
17873    xs: &[f32],
17874    b: usize,
17875    hidden: usize,
17876    pool: Option<&Pool>,
17877) -> Vec<f32> {
17878    let t0 = std::time::Instant::now();
17879    let routes: Vec<_> = xs
17880        .chunks_exact(hidden)
17881        .map(|x| moe_ffn_route(m, x, pool, None))
17882        .collect();
17883    slot.note_route(t0.elapsed().as_nanos() as u64);
17884    if let Some(out) = slot.forward_rows(li, m, xs, &routes, pool) {
17885        return out;
17886    }
17887    let mut out = Vec::with_capacity(b * hidden);
17888    for (x, r) in xs.chunks_exact(hidden).zip(&routes) {
17889        let row = slot.forward(li, m, x, r, pool).unwrap_or_else(|| {
17890            // A failed bank must not stream missing experts into the arena.
17891            crate::qtensor::float_activations_scope(|| {
17892                crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, r, pool))
17893            })
17894        });
17895        out.extend(row);
17896    }
17897    out
17898}
17899
17900/// One-shot report of whether the whole-token wgpu graph actually formed.
17901/// A refusal silently reverts to the per-op path, which is how a model can
17902/// look "GPU-accelerated" while every layer walks the host.  A device prefix
17903/// is tracked separately because it still pays a host boundary for the tail.
17904fn graph_note(built: bool, layers_run: usize, total_layers: usize) {
17905    use std::sync::atomic::{AtomicBool, Ordering};
17906    if built {
17907        GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
17908        if total_layers > 0 && layers_run < total_layers {
17909            GRAPH_TOK_PREFIX.fetch_add(1, Ordering::Relaxed);
17910        } else {
17911            GRAPH_TOK_FULL.fetch_add(1, Ordering::Relaxed);
17912        }
17913    } else {
17914        GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
17915    }
17916    static SAID: AtomicBool = AtomicBool::new(false);
17917    if !SAID.swap(true, Ordering::Relaxed) {
17918        if built {
17919            tracing::info!("wgpu whole-token graph: ACTIVE");
17920        } else {
17921            tracing::warn!("wgpu whole-token graph refused — per-op path");
17922        }
17923    }
17924}
17925
17926/// Whole-token graph outcomes, process-wide: a benchmark that claims a
17927/// GPU number while MISS climbs is measuring the CPU — the honest-bench
17928/// contract makes that an error, not a footnote.
17929pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17930pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17931/// Graph calls that returned a hidden after running only a leading device
17932/// prefix.  These are valid hybrid executions but must not be reported as a
17933/// full GPU graph in benchmark evidence.
17934pub static GRAPH_TOK_PREFIX: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17935/// Graph calls that covered the complete requested layer span.
17936pub static GRAPH_TOK_FULL: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
17937
17938/// Native Metal TokenGraph completion counters. These are incremented only
17939/// after checked command-buffer completion and successful readback, so a
17940/// fused-head NLL report can prove the route rather than infer it from env.
17941pub static METAL_GRAPH_TOK_OK: std::sync::atomic::AtomicU64 =
17942    std::sync::atomic::AtomicU64::new(0);
17943pub static METAL_GRAPH_HEAD_OK: std::sync::atomic::AtomicU64 =
17944    std::sync::atomic::AtomicU64::new(0);
17945pub static METAL_GRAPH_HEAD_MISS: std::sync::atomic::AtomicU64 =
17946    std::sync::atomic::AtomicU64::new(0);
17947pub static METAL_GRAPH_LAYERS: std::sync::atomic::AtomicU64 =
17948    std::sync::atomic::AtomicU64::new(0);
17949pub static METAL_GRAPH_ERRORS: std::sync::atomic::AtomicU64 =
17950    std::sync::atomic::AtomicU64::new(0);
17951/// Ordinary native-Metal rows-prefill admissions and completed rows.  These
17952/// counters are separate from TokenGraph token/head counts so a batch NLL
17953/// receipt cannot accidentally claim serial execution as batched.
17954pub static METAL_PREFILL_CHUNKS: std::sync::atomic::AtomicU64 =
17955    std::sync::atomic::AtomicU64::new(0);
17956pub static METAL_PREFILL_ROWS: std::sync::atomic::AtomicU64 =
17957    std::sync::atomic::AtomicU64::new(0);
17958pub static METAL_PREFILL_HEAD_ROWS: std::sync::atomic::AtomicU64 =
17959    std::sync::atomic::AtomicU64::new(0);
17960pub static METAL_PREFILL_ERRORS: std::sync::atomic::AtomicU64 =
17961    std::sync::atomic::AtomicU64::new(0);
17962
17963/// `CMF_MOE_BATCH=0` restores the per-expert serial loop — the A/B lever
17964/// for the batched kernel, and how its bit-identity is checked.
17965fn moe_batch_enabled() -> bool {
17966    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17967    *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
17968}
17969
17970/// Two-dispatch CPU MoE: every routed expert (and the shared one) fused
17971/// into one gate/up/SiLU dispatch and one down dispatch, instead of two
17972/// pool barriers per expert. Bit-identical to the serial loop below —
17973/// see `moe_gate_up_many` / `moe_down_many`. `None` = the batched kernel
17974/// does not cover this layer, walk the serial path.
17975fn moe_ffn_cpu_batched(
17976    m: &MoeFfn,
17977    x: &[f32],
17978    idx: &[usize],
17979    p: &[f32],
17980    wsum: f32,
17981    pool: Option<&Pool>,
17982) -> Option<Vec<f32>> {
17983    if idx.is_empty() || !moe_batch_enabled() {
17984        return None;
17985    }
17986    // The bake probe reads per-neuron activation mass out of the
17987    // single-expert path; batching would skip it. Rare and offline —
17988    // hand those runs to the serial loop.
17989    if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
17990        return None;
17991    }
17992    let n = idx.len() + usize::from(m.shared.is_some());
17993    let mut pairs = Vec::with_capacity(n);
17994    let mut downs = Vec::with_capacity(n);
17995    let mut ws = Vec::with_capacity(n);
17996    for &e in idx {
17997        let d = &m.experts[e];
17998        if d.act != Act::Silu {
17999            return None;
18000        }
18001        pairs.push((&d.gate_proj, &d.up_proj));
18002        downs.push(&d.down_proj);
18003        ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
18004    }
18005    // The shared expert goes last, matching the serial loop's order —
18006    // the f32 accumulation order is part of the bit-identity claim.
18007    if let Some((se, gate)) = &m.shared {
18008        if se.act != Act::Silu {
18009            return None;
18010        }
18011        let g = gate.as_ref().map_or(1.0, |gate| {
18012            let mut gl = [0.0f32; 1];
18013            gate.matvec(x, &mut gl, pool);
18014            1.0 / (1.0 + (-gl[0]).exp())
18015        });
18016        pairs.push((&se.gate_proj, &se.up_proj));
18017        downs.push(&se.down_proj);
18018        ws.push(g);
18019    }
18020    let inter = pairs[0].0.rows();
18021    let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
18022    if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
18023        return None;
18024    }
18025    let mut out = attention::take_buf(x.len());
18026    if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
18027        attention::recycle_buf(&mut out);
18028        return None;
18029    }
18030    Some(out)
18031}
18032
18033/// Exact CPU completion for the routed experts a dynamic device cache did
18034/// not contain. The weights are already the router's final normalized mix.
18035/// Keeping this independent of `MoeFfn` makes the job `Sync`: its routing
18036/// statistics live in a `RefCell`, while the immutable expert tensors can be
18037/// evaluated safely in parallel with the GPU's resident subset.
18038pub(crate) fn moe_cold_experts_cpu(
18039    experts: &[(&DenseFfn, f32)],
18040    x: &[f32],
18041    pool: Option<&Pool>,
18042) -> Vec<f32> {
18043    let mut out = attention::take_buf(x.len());
18044    if experts.is_empty() {
18045        return out;
18046    }
18047    let pairs: Vec<_> = experts
18048        .iter()
18049        .map(|(e, _)| (&e.gate_proj, &e.up_proj))
18050        .collect();
18051    let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
18052    let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
18053    let inter = experts[0].0.gate_proj.rows();
18054    let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
18055    if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
18056        && QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
18057    {
18058        return out;
18059    }
18060    out.fill(0.0);
18061    for &(expert, weight) in experts {
18062        let mut one = dense_ffn(expert, x, pool);
18063        for (o, v) in out.iter_mut().zip(&one) {
18064            *o += weight * v;
18065        }
18066        attention::recycle_buf(&mut one);
18067    }
18068    out
18069}
18070
18071/// Cold part of a short bank batch. Share each expert's weight stream
18072/// across its tokens, but reduce contributions in each token's route order.
18073/// On an unsupported CPU/layout, retain the single-token cold kernels.
18074pub(crate) fn moe_cold_experts_rows_cpu(
18075    jobs: &[Vec<(&DenseFfn, f32)>],
18076    xs: &[f32],
18077    hidden: usize,
18078    pool: Option<&Pool>,
18079) -> Vec<f32> {
18080    let mut out = vec![0.0; xs.len()];
18081    let mut experts: Vec<&DenseFfn> = Vec::new();
18082    let mut groups: Vec<Vec<usize>> = Vec::new();
18083    let mut terms = vec![Vec::new(); jobs.len()];
18084    for (r, row) in jobs.iter().enumerate() {
18085        for &(e, w) in row {
18086            let g = match experts.iter().position(|&d| std::ptr::eq(d, e)) {
18087                Some(g) => g,
18088                None => {
18089                    experts.push(e);
18090                    groups.push(Vec::new());
18091                    groups.len() - 1
18092                }
18093            };
18094            terms[r].push((g, groups[g].len(), w));
18095            groups[g].push(r);
18096        }
18097    }
18098    if experts.is_empty() {
18099        return out;
18100    }
18101    let pairs: Vec<_> = experts.iter().map(|e| (&e.gate_proj, &e.up_proj)).collect();
18102    let downs: Vec<_> = experts.iter().map(|e| &e.down_proj).collect();
18103    let lens: Vec<_> = groups.iter().map(Vec::len).collect();
18104    let count: usize = lens.iter().sum();
18105    let mut acts = vec![vec![0.0; experts[0].gate_proj.rows()]; count];
18106    let mut ds = vec![vec![0.0; hidden]; count];
18107    if QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut acts, pool)
18108        && QTensor::moe_down_rows(&downs, &lens, &acts, &mut ds, pool)
18109    {
18110        let mut offset = 0;
18111        let offsets: Vec<_> = lens
18112            .iter()
18113            .map(|&n| {
18114                let start = offset;
18115                offset += n;
18116                start
18117            })
18118            .collect();
18119        for (r, terms) in terms.iter().enumerate() {
18120            for &(g, slot, w) in terms {
18121                for (o, &v) in out[r * hidden..(r + 1) * hidden]
18122                    .iter_mut()
18123                    .zip(&ds[offsets[g] + slot])
18124                {
18125                    *o += w * v;
18126                }
18127            }
18128        }
18129    } else {
18130        for (r, jobs) in jobs.iter().enumerate() {
18131            if !jobs.is_empty() {
18132                let mut row = moe_cold_experts_cpu(jobs, &xs[r * hidden..(r + 1) * hidden], pool);
18133                out[r * hidden..(r + 1) * hidden].copy_from_slice(&row);
18134                attention::recycle_buf(&mut row);
18135            }
18136        }
18137    }
18138    out
18139}
18140
18141/// The pure-CPU MoE expert loop (also the fallback of every GPU refusal).
18142fn moe_ffn_cpu(
18143    m: &MoeFfn,
18144    x: &[f32],
18145    idx: &[usize],
18146    p: &[f32],
18147    wsum: f32,
18148    pool: Option<&Pool>,
18149) -> Vec<f32> {
18150    if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
18151        return out;
18152    }
18153    let mut out = attention::take_buf(x.len());
18154    for &e in idx {
18155        let mut eo = dense_ffn(&m.experts[e], x, pool);
18156        let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
18157        for i in 0..out.len() {
18158            out[i] += w * eo[i];
18159        }
18160        attention::recycle_buf(&mut eo);
18161    }
18162    if let Some((se, gate)) = &m.shared {
18163        let mut so = dense_ffn(se, x, pool);
18164        let g = gate.as_ref().map_or(1.0, |gate| {
18165            let mut gl = [0.0f32; 1];
18166            gate.matvec(x, &mut gl, pool);
18167            1.0 / (1.0 + (-gl[0]).exp())
18168        });
18169        for i in 0..out.len() {
18170            out[i] += g * so[i];
18171        }
18172        attention::recycle_buf(&mut so);
18173    }
18174    out
18175}
18176
18177/// DeepSeek-V2 MLA forward, expand-to-MHA form (see `AttnKind::Mla`):
18178/// per token the latent expands to every head's K/V and the ordinary
18179/// cache + grouped attend do the rest. K head layout is [rope | nope]
18180/// (rotary_dim = qk_rope rotates the shared rope key and each q head's
18181/// prefix); V rows are zero-padded to the K head_dim inside the cache
18182/// and the pad is sliced off before O. Attention importance is not
18183/// accumulated for MLA yet (no eviction interplay).
18184#[allow(clippy::too_many_arguments)]
18185pub(crate) fn mla_attention(
18186    w: &MlaWeights,
18187    normed: &[f32],
18188    cache: &mut crate::kv_cache::LayerKvCache,
18189    position: usize,
18190    inv_freq: &[f32],
18191    rope_scale: f32,
18192    eps: f64,
18193    pool: Option<&Pool>,
18194) -> Vec<f32> {
18195    let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
18196    let hd = dr + dn;
18197    let mut q = vec![0.0f32; nh * hd];
18198    match (&w.q_a, &w.q_a_norm) {
18199        (Some(qa), Some(qn)) => {
18200            let mut t = vec![0.0f32; qa.rows()];
18201            qa.matvec(normed, &mut t, pool);
18202            let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
18203            w.q_proj.matvec(&tn, &mut q, pool);
18204        }
18205        _ => w.q_proj.matvec(normed, &mut q, pool),
18206    }
18207    let mut ca = vec![0.0f32; lora + dr];
18208    w.kv_a.matvec(normed, &mut ca, pool);
18209    let (c_lat, k_rope) = ca.split_at_mut(lora);
18210    let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
18211    let mut kvb = vec![0.0f32; nh * (dn + dv)];
18212    w.kv_b.matvec(&latn, &mut kvb, pool);
18213    if !w.nope {
18214        attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
18215    }
18216    for h in 0..nh {
18217        if !w.nope {
18218            attention::rope_rotate_scaled(
18219                &mut q[h * hd..h * hd + dr],
18220                position,
18221                inv_freq,
18222                rope_scale,
18223            );
18224        }
18225    }
18226    let mut k = vec![0.0f32; nh * hd];
18227    let mut v = vec![0.0f32; nh * hd];
18228    for h in 0..nh {
18229        k[h * hd..h * hd + dr].copy_from_slice(k_rope);
18230        k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
18231        v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
18232    }
18233    cache.append(&k, &v, &vec![true; nh]);
18234    let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
18235    attention::recycle_buf(&mut imp);
18236    let mut ov = vec![0.0f32; nh * dv];
18237    for h in 0..nh {
18238        ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
18239    }
18240    let mut out = vec![0.0f32; w.o_proj.rows()];
18241    w.o_proj.matvec(&ov, &mut out, pool);
18242    out
18243}
18244
18245/// Gemma-4 dual-branch FFN (spec: see `FfnKind::DenseMoe`). The dense
18246/// branch reads the pre-FFN-normed activation; the router and the
18247/// expert branch read the RAW residual — the router through a
18248/// scale-less rms norm (its constant gain is folded into the weights),
18249/// the experts through `pre_norm_2`. CPU path; GPU graphs refuse the
18250/// layer kind honestly.
18251fn dense_moe_ffn(
18252    dm: &DenseMoeFfn,
18253    x_normed: &[f32],
18254    h_raw: &[f32],
18255    eps: f64,
18256    norm_style: NormStyle,
18257    pool: Option<&Pool>,
18258) -> Vec<f32> {
18259    let mut d = dense_ffn(&dm.dense, x_normed, pool);
18260    d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
18261    let m = &dm.moe;
18262    let ne = m.experts.len();
18263    let mut logits = vec![0.0f32; ne];
18264    if m.router_input_norm {
18265        let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
18266        let inv = 1.0 / (ss + eps as f32).sqrt();
18267        let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
18268        m.router.matvec(&xr, &mut logits, pool);
18269    } else {
18270        m.router.matvec(h_raw, &mut logits, pool);
18271    }
18272    let (idx, p, wsum) = moe_route(&logits, m, None);
18273    {
18274        let mut st = m.stats.borrow_mut();
18275        if st.len() < ne {
18276            st.resize(ne, 0);
18277        }
18278        for &e in &idx {
18279            st[e] += 1;
18280        }
18281    }
18282    let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
18283    let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
18284    let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
18285    for (di, mi) in d.iter_mut().zip(&mo) {
18286        *di += mi;
18287    }
18288    d
18289}
18290
18291/// Building the MoE-layer GPU jobs: all selected experts (+shared) must
18292/// be q8_2f-Mapped from the primary mapping; otherwise None → CPU path.
18293/// One-shot report of why the MoE GPU block refused. A silent `?` here
18294/// sends every expert to the CPU with nothing in the logs to say so —
18295/// which is exactly how a q4tp MoE model looked "GPU-accelerated" while
18296/// running entirely on the host.
18297fn moe_gpu_refused(why: &'static str) {
18298    use std::sync::atomic::{AtomicBool, Ordering};
18299    static SAID: AtomicBool = AtomicBool::new(false);
18300    if !SAID.swap(true, Ordering::Relaxed) {
18301        tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
18302    }
18303}
18304
18305fn moe_ffn_gpu(
18306    m: &MoeFfn,
18307    x: &[f32],
18308    idx: &[usize],
18309    p: &[f32],
18310    wsum: f32,
18311    pool: Option<&Pool>,
18312) -> Option<Vec<f32>> {
18313    use crate::gpu::MoeJob;
18314
18315    let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
18316    let mut model_ref = None;
18317    for &e in idx {
18318        if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
18319            moe_gpu_refused("push_job(expert)");
18320            return None;
18321        }
18322    }
18323    if let Some((se, gate)) = &m.shared {
18324        let g = gate.as_ref().map_or(1.0, |gate| {
18325            let mut gl = [0.0f32; 1];
18326            gate.matvec(x, &mut gl, pool);
18327            1.0 / (1.0 + (-gl[0]).exp())
18328        });
18329        if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
18330            moe_gpu_refused("push_job(shared)");
18331            return None;
18332        }
18333    }
18334    let Some(model) = model_ref else {
18335        moe_gpu_refused("no model_ref");
18336        return None;
18337    };
18338    let hidden = jobs[0].down.1;
18339    let mut out = vec![0.0f32; hidden];
18340    if crate::gpu::moe_block(&model, &jobs, &mut out) {
18341        Some(out)
18342    } else {
18343        moe_gpu_refused("gpu::moe_block");
18344        None
18345    }
18346}
18347
18348/// Single-position FFN dispatch.
18349fn ffn_forward(
18350    ffn: &FfnKind,
18351    x: &[f32],
18352    pool: Option<&Pool>,
18353    experts_allowed: Option<&[bool]>,
18354) -> Vec<f32> {
18355    match ffn {
18356        FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
18357        FfnKind::Dense(d) => dense_ffn(d, x, pool),
18358        FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
18359        // Dual-branch layers need the raw residual — their callers
18360        // dispatch dense_moe_ffn directly; the auxiliary paths that land
18361        // here (MTP draft, o1 replay) do not co-occur with gemma-4 MoE.
18362        FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18363    }
18364}
18365
18366/// Fused two-position FFN: gate/up/down streamed once (dense). MoE
18367/// falls back to two singles — expert sets differ per position, there
18368/// is nothing to fuse.
18369fn ffn_forward_pair(
18370    ffn: &FfnKind,
18371    x1: &[f32],
18372    x2: &[f32],
18373    pool: Option<&Pool>,
18374    experts_allowed: Option<&[bool]>,
18375) -> (Vec<f32>, Vec<f32>) {
18376    let d = match ffn {
18377        // A tube layer has nothing to fuse across the pair — the tubes
18378        // are separate matrices; two singles are the honest path.
18379        FfnKind::Dense(d) if !d.segs.is_empty() => {
18380            return (
18381                tube_ffn(d, x1, 1, pool, None),
18382                tube_ffn(d, x2, 1, pool, None),
18383            );
18384        }
18385        FfnKind::Dense(d) => d,
18386        FfnKind::Moe(m) => {
18387            return (
18388                moe_ffn(m, x1, pool, experts_allowed),
18389                moe_ffn(m, x2, pool, experts_allowed),
18390            );
18391        }
18392        FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18393    };
18394    let inter = d.gate_proj.rows();
18395    FFN_SCRATCH.with(|s| {
18396        let mut s = s.borrow_mut();
18397        let [g1, g2, u1, u2] = &mut *s;
18398        g1.resize(inter, 0.0);
18399        g2.resize(inter, 0.0);
18400        u1.resize(inter, 0.0);
18401        u2.resize(inter, 0.0);
18402        // Multi-matrix pair job: gate+up under one pool dispatch
18403        // (o1s = lane-1 outputs across tensors, o2s = lane-2).
18404        QTensor::matvec2_many(
18405            [&d.gate_proj, &d.up_proj],
18406            x1,
18407            x2,
18408            [g1.as_mut_slice(), u1.as_mut_slice()],
18409            [g2.as_mut_slice(), u2.as_mut_slice()],
18410            pool,
18411        );
18412        for i in 0..inter {
18413            g1[i] = d.act.combine(g1[i], u1[i]);
18414            g2[i] = d.act.combine(g2[i], u2[i]);
18415        }
18416        let mut o1 = attention::take_buf(d.down_proj.rows());
18417        let mut o2 = attention::take_buf(d.down_proj.rows());
18418        d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
18419        (o1, o2)
18420    })
18421}
18422
18423#[cfg(test)]
18424mod tests {
18425
18426    /// The 0.7.6 prefill-chunk rule: a plain dense stack wholly on a
18427    /// discrete card reads the prompt in wide chunks on x86; every other
18428    /// case keeps the width it had (the GDN-hybrid, MoE and DeepSeek paths
18429    /// were tuned on hardware not measured for this change).
18430    #[test]
18431    fn prefill_chunk_rule_widens_only_dense_on_discrete() {
18432        use super::{
18433            prefill_chunk_rule, ChunkHost, ChunkStackFacts, DISCRETE_DENSE_PREFILL_CHUNK,
18434        };
18435        let dense_card = ChunkStackFacts {
18436            plain_dense: true,
18437            discrete: true,
18438            gpu_on: true,
18439            ..Default::default()
18440        };
18441        assert!(dense_card.dense_on_discrete());
18442        // The bug: a dense Llama on a Vulkan RTX 3090 got 48.
18443        assert_eq!(
18444            prefill_chunk_rule(None, ChunkHost::Other, dense_card.dense_on_discrete()),
18445            DISCRETE_DENSE_PREFILL_CHUNK
18446        );
18447        assert!(DISCRETE_DENSE_PREFILL_CHUNK > 48);
18448        for (label, facts) in [
18449            ("GDN hybrid / MoE / DeepSeek stack", ChunkStackFacts { plain_dense: false, ..dense_card }),
18450            ("integrated GPU", ChunkStackFacts { discrete: false, ..dense_card }),
18451            ("CPU only", ChunkStackFacts { gpu_on: false, discrete: false, ..dense_card }),
18452            ("capacity split", ChunkStackFacts { capacity_split: true, ..dense_card }),
18453            ("multi-GPU plan", ChunkStackFacts { multi_gpu: true, ..dense_card }),
18454            ("O(1) layers", ChunkStackFacts { o1: true, ..dense_card }),
18455        ] {
18456            assert!(!facts.dense_on_discrete(), "{label}");
18457            assert_eq!(
18458                prefill_chunk_rule(None, ChunkHost::Other, facts.dense_on_discrete()),
18459                48,
18460                "{label} keeps the historical x86 chunk"
18461            );
18462        }
18463        // Other hosts are untouched whatever the model.
18464        for dense in [false, true] {
18465            assert_eq!(prefill_chunk_rule(None, ChunkHost::Macos, dense), 512);
18466            assert_eq!(prefill_chunk_rule(None, ChunkHost::Aarch64, dense), 256);
18467        }
18468        // CMF_PREFILL_CHUNK still wins everywhere (and is clamped to ≥ 1).
18469        for host in [ChunkHost::Macos, ChunkHost::Aarch64, ChunkHost::Other] {
18470            for dense in [false, true] {
18471                assert_eq!(prefill_chunk_rule(Some(48), host, dense), 48);
18472                assert_eq!(prefill_chunk_rule(Some(0), host, dense), 1);
18473            }
18474        }
18475    }
18476
18477    #[test]
18478    fn kv_reuse_plan_pulls_rows_decode_wrote_only_on_the_device() {
18479        use super::{ReuseLayer, ReusePlan, kv_reuse_plan};
18480        let full = |host_rows, device_rows| ReuseLayer {
18481            full: true,
18482            host_rows,
18483            device_rows,
18484            device_state: false,
18485        };
18486        // Turn 1: 300-token prompt prefilled on the host, 40 tokens decoded
18487        // by the wgpu graph into the device mirror only. Turn 2 reuses 339.
18488        assert_eq!(
18489            kv_reuse_plan(339, &[full(300, Some(339)), full(300, Some(339))]),
18490            ReusePlan::Pull(vec![(0, 300, 339), (1, 300, 339)])
18491        );
18492        // CPU / Metal: the host owner already holds every forwarded row.
18493        assert_eq!(kv_reuse_plan(339, &[full(339, None)]), ReusePlan::Ready);
18494        // A mirror past the prefix is fine for the host (it gets rewound).
18495        assert_eq!(kv_reuse_plan(339, &[full(339, Some(345))]), ReusePlan::Ready);
18496        // GPU prefix / CPU tail: only the device layers lag.
18497        assert_eq!(
18498            kv_reuse_plan(339, &[full(300, Some(339)), full(339, None)]),
18499            ReusePlan::Pull(vec![(0, 300, 339)])
18500        );
18501        // The device cannot supply the missing rows: never continue.
18502        assert_eq!(kv_reuse_plan(339, &[full(300, Some(320))]), ReusePlan::Fresh);
18503        assert_eq!(kv_reuse_plan(339, &[full(300, None)]), ReusePlan::Fresh);
18504        assert_eq!(kv_reuse_plan(339, &[full(350, None)]), ReusePlan::Fresh);
18505        // A recurrent state advanced on the device cannot be handed to a
18506        // host prefill (it is not rewindable and the host copy is stale).
18507        let conv = |device_state| ReuseLayer {
18508            full: false,
18509            host_rows: 0,
18510            device_rows: None,
18511            device_state,
18512        };
18513        assert_eq!(
18514            kv_reuse_plan(339, &[conv(true), full(300, Some(339))]),
18515            ReusePlan::Fresh
18516        );
18517        assert_eq!(kv_reuse_plan(339, &[conv(false), full(339, None)]), ReusePlan::Ready);
18518    }
18519
18520    #[test]
18521    fn nll_graph_policy_scopes_only_the_fused_head() {
18522        for (label, unmasked, prefer_graph, native_metal, want_graph, want_head) in [
18523            // A Vulkan/Wgpu hidden-only graph remains the quality route.
18524            ("vulkan graph", true, true, false, true, false),
18525            // Native Metal adds the strict fused graph-head contract.
18526            ("native Metal graph", true, true, true, true, true),
18527            // Masked NLL and the explicit non-graph fallback remain unchanged.
18528            ("masked", false, true, false, false, false),
18529            ("graph disabled", true, false, true, false, false),
18530        ] {
18531            let (graph_quality, graph_head_required) =
18532                super::nll_graph_policy(unmasked, prefer_graph, native_metal);
18533            assert_eq!(graph_quality, want_graph, "{label}: graph quality");
18534            assert_eq!(graph_head_required, want_head, "{label}: fused head");
18535        }
18536    }
18537
18538    #[test]
18539    fn mtp_prefill_pair_boundaries_skip_only_final_prompt_row() {
18540        assert_eq!(mtp_prefill_pair_count(0, 128, 256), 128);
18541        assert_eq!(mtp_prefill_pair_count(128, 256, 256), 127);
18542        assert_eq!(mtp_prefill_pair_count(0, 256, 256), 255);
18543        assert_eq!(mtp_prefill_pair_count(256, 256, 256), 0);
18544        assert_eq!(mtp_prefill_pair_count(300, 320, 256), 0);
18545    }
18546
18547    #[test]
18548    fn cancel_flag_stops_generation() {
18549        let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
18550        // Set before the call: the prefill loops honour it, the run
18551        // returns immediately with the cancelled reason and no tokens.
18552        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
18553        let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
18554        assert_eq!(r.finish_reason, "cancelled");
18555        assert!(
18556            r.token_ids.is_empty(),
18557            "no tokens after cancel: {:?}",
18558            r.token_ids
18559        );
18560        assert_eq!(p.kv_cache.seq_len(), 0);
18561        assert!(p.kv_history.is_empty());
18562        assert!(!p.graph_want_logits);
18563        assert!(p.graph_logits.is_none());
18564        // Flag auto-cleared: the next call generates normally.
18565        let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
18566        assert_ne!(r2.finish_reason, "cancelled");
18567    }
18568    use super::*;
18569
18570    /// sparse_ffn_quant must equal a dense FFN where inactive neurons are
18571    /// zeroed (mask × mmap correctness). On F32 tensors this is EXACT —
18572    /// it validates the row_dot / add_col_scaled / scatter indexing, the
18573    /// bug-prone part. The q8 branches reuse the golden-tested linear
18574    /// The per-token sparse path reads a transposed `down`; it must
18575    /// agree with the arm that computes everything and zeroes the
18576    /// losers, or the speed measurement is measuring a different model.
18577    #[test]
18578    fn dynamic_ffn_equals_the_zeroing_arm() {
18579        let (hidden, inter) = (8usize, 32usize);
18580        let synth = |n: usize, salt: usize| -> Vec<f32> {
18581            (0..n)
18582                .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
18583                .collect()
18584        };
18585        let down = synth(hidden * inter, 3);
18586        let mut down_t = vec![0.0f32; inter * hidden];
18587        for r in 0..hidden {
18588            for c in 0..inter {
18589                down_t[c * hidden + r] = down[r * inter + c];
18590            }
18591        }
18592        let d = DenseFfn {
18593            gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18594            up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18595            down_proj: QTensor::from_f32(down.clone(), hidden, inter),
18596            act: Act::Silu,
18597            down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
18598            segs: Vec::new(),
18599        };
18600        let x = synth(hidden, 11);
18601        let k = 12usize;
18602        let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
18603        // Reference: full compute, keep the k loudest |silu(gate)|.
18604        let mut g = vec![0.0f32; inter];
18605        d.gate_proj.matvec(&x, &mut g, None);
18606        let mut u = vec![0.0f32; inter];
18607        d.up_proj.matvec(&x, &mut u, None);
18608        for v in g.iter_mut() {
18609            *v = inference::silu(*v);
18610        }
18611        keep_top_k(&mut g, k);
18612        for i in 0..inter {
18613            g[i] *= u[i];
18614        }
18615        let mut want = vec![0.0f32; hidden];
18616        d.down_proj.matvec(&g, &mut want, None);
18617        for (a, b) in want.iter().zip(&got) {
18618            assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
18619        }
18620    }
18621
18622    /// A tube layer is the same layer, re-cut. With every tube open the
18623    /// answer must equal the dense FFN over the concatenated neurons
18624    /// (the permutation is an identity on the layer's function); with a
18625    /// tube closed it must equal the dense FFN with those neurons
18626    /// zeroed — the mask semantics, now paid for in bytes not read.
18627    #[test]
18628    fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
18629        let (hidden, core, tube) = (8usize, 12usize, 8usize);
18630        let inter = core + tube;
18631        let synth = |n: usize, salt: usize| -> Vec<f32> {
18632            (0..n)
18633                .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
18634                .collect()
18635        };
18636        let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
18637        let d_all = synth(hidden * inter, 3);
18638        // The dense layer, and the same weights cut into core + tube.
18639        let dense = DenseFfn {
18640            gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
18641            up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
18642            down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
18643            act: Act::Silu,
18644            down_t: None,
18645            segs: Vec::new(),
18646        };
18647        let rows =
18648            |v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
18649        let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
18650            let mut o = Vec::with_capacity(hidden * (b - a));
18651            for r in 0..hidden {
18652                o.extend_from_slice(&v[r * inter + a..r * inter + b]);
18653            }
18654            o
18655        };
18656        let tubed = DenseFfn {
18657            down_t: None,
18658            gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
18659            up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
18660            down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
18661            act: Act::Silu,
18662            segs: vec![FfnSeg {
18663                gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
18664                up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
18665                down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
18666                start: core,
18667                width: tube,
18668            }],
18669        };
18670        let x = synth(hidden, 7);
18671        let want = dense_ffn(&dense, &x, None);
18672        let got = tube_ffn(&tubed, &x, 1, None, None);
18673        for (a, b) in want.iter().zip(&got) {
18674            assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
18675        }
18676        // Closed tube: bits on for the core, off for the tube.
18677        let mut bits = vec![0u8; inter.div_ceil(8)];
18678        for n in 0..core {
18679            bits[n / 8] |= 1 << (n % 8);
18680        }
18681        let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18682        let masked = dense_ffn_masked(&dense, &x, None, &bits);
18683        for (a, b) in masked.iter().zip(&closed) {
18684            assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
18685        }
18686        // The batched arm must agree with the single-position one.
18687        let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
18688        for (a, b) in closed.iter().zip(&batch) {
18689            assert_eq!(a, b, "batch arm disagrees with decode arm");
18690        }
18691    }
18692
18693    /// scale, structurally identical to the matvec kernels.
18694    #[test]
18695    fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
18696        let (hidden, inter) = (16usize, 40usize);
18697        let synth = |n: usize, salt: usize| -> Vec<f32> {
18698            (0..n)
18699                .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
18700                .collect()
18701        };
18702        let d = DenseFfn {
18703            gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
18704            up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
18705            down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
18706            act: Act::Silu,
18707            down_t: None,
18708            segs: Vec::new(),
18709        };
18710        let x = synth(hidden, 9);
18711        // Active = every 3rd neuron.
18712        let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
18713
18714        let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
18715
18716        // Reference: full dense FFN but g[i]=0 for inactive neurons.
18717        let mut g = vec![0.0f32; inter];
18718        d.gate_proj.matvec(&x, &mut g, None);
18719        let mut u = vec![0.0f32; inter];
18720        d.up_proj.matvec(&x, &mut u, None);
18721        let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
18722        for i in 0..inter {
18723            g[i] = if act_set.contains(&(i as u16)) {
18724                inference::silu(g[i]) * u[i]
18725            } else {
18726                0.0
18727            };
18728        }
18729        let mut reference = vec![0.0f32; hidden];
18730        d.down_proj.matvec(&g, &mut reference, None);
18731
18732        let max_d = sparse
18733            .iter()
18734            .zip(&reference)
18735            .map(|(a, b)| (a - b).abs())
18736            .fold(0.0f32, f32::max);
18737        assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
18738    }
18739
18740    /// Attach a synthetic MTP head (same structure as a main layer).
18741    fn attach_test_mtp(p: &mut Pipeline) {
18742        let (h, inter, heads, kv, hd) = (
18743            p.hidden_size,
18744            p.intermediate_size,
18745            p.num_heads,
18746            p.num_kv_heads,
18747            p.head_dim,
18748        );
18749        let synth = |n: usize, salt: usize| -> Vec<f32> {
18750            (0..n)
18751                .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
18752                .collect()
18753        };
18754        let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
18755            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
18756        };
18757        p.mtp = Some(MtpModule {
18758            enorm: vec![1.0; h],
18759            hnorm: vec![1.0; h],
18760            eh_proj: qt(h, 2 * h, 301),
18761            layer: LayerWeights {
18762                input_norm: vec![1.0; h],
18763                post_norm: vec![1.0; h],
18764                attn_out_norm: None,
18765                ffn_out_norm: None,
18766                layer_scale: None,
18767                ffn: FfnKind::Dense(DenseFfn {
18768                    gate_proj: qt(inter, h, 315),
18769                    up_proj: qt(inter, h, 316),
18770                    down_proj: qt(h, inter, 317),
18771                    act: Act::Silu,
18772                    down_t: None,
18773                    segs: Vec::new(),
18774                }),
18775                attn: AttnKind::Full {
18776                    bias: None,
18777                    wq: qt(heads * hd, h, 311),
18778                    wk: qt(kv * hd, h, 312),
18779                    wv: qt(kv * hd, h, 313),
18780                    wo: qt(h, heads * hd, 314),
18781                    q_norm: None,
18782                    k_norm: None,
18783                    output_gate: false,
18784                    softplus_gate: None,
18785                },
18786            },
18787            final_norm: vec![1.0; h],
18788            kv: crate::kv_cache::LayerKvCache::new(kv, hd),
18789        });
18790    }
18791
18792    #[test]
18793    fn speculative_equals_vanilla_greedy() {
18794        // Speculative decode and the wgpu token graph are mutually
18795        // exclusive; a leaked CMF_GPU=wgpu from a parallel gpu test
18796        // would silently disable drafting. Pin the graph off.
18797        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
18798        let run = |spec: bool| {
18799            let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18800            p.sampler_config.temperature = 0.0;
18801            attach_test_mtp(&mut p);
18802            p.speculative = spec;
18803            let r = p.generate("abcdef", 12, None, None).unwrap();
18804            (r.token_ids, r.mtp_drafted, r.mtp_accepted)
18805        };
18806        let (vanilla, d0, _) = run(false);
18807        let (spec, d1, a1) = run(true);
18808        assert_eq!(d0, 0, "vanilla path must not draft");
18809        assert!(d1 > 0, "speculative path must draft");
18810        assert_eq!(
18811            vanilla, spec,
18812            "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
18813        );
18814    }
18815
18816    #[test]
18817    fn speculative_accepts_constant_oracle() {
18818        // See speculative_equals_vanilla_greedy: pin the wgpu graph off.
18819        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
18820        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18821        p.sampler_config.temperature = 0.0;
18822        p.sampler_config.repetition_penalty = 1.0;
18823        // Constant lm_head → every logit equal → both the main model and
18824        // the draft head argmax to token 0: acceptance must be 100%.
18825        p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
18826        attach_test_mtp(&mut p);
18827        p.speculative = true;
18828        let r = p.generate("abcd", 10, None, None).unwrap();
18829        assert!(r.mtp_drafted > 0);
18830        assert_eq!(
18831            r.mtp_accepted, r.mtp_drafted,
18832            "constant logits → every draft accepted"
18833        );
18834        // Ties resolve to the same token in both the main and draft
18835        // heads — the sequence is one repeated token.
18836        assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
18837    }
18838
18839    #[test]
18840    fn empty_prompt_is_an_error_not_a_panic() {
18841        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
18842        let r = p.generate("", 4, None, None);
18843        assert!(r.is_err(), "empty prompt must be a clean error");
18844    }
18845
18846    #[test]
18847    fn every_token_enters_kv_exactly_once() {
18848        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18849        // Greedy so no RNG variance; byte tokenizer → 3 prompt tokens.
18850        p.sampler_config.temperature = 0.0;
18851        let r = p.generate("abc", 2, None, None).unwrap();
18852        assert_eq!(r.prompt_tokens, 3);
18853        // prompt(3) + first sampled token forwarded before second logits:
18854        // step0 samples from prefill hidden (no extra forward), then
18855        // forwards t1 → cache 4; step1 samples, loop ends (max_tokens).
18856        assert_eq!(
18857            p.kv_cache.seq_len(),
18858            3 + r.tokens_generated - 1,
18859            "each token must be cached exactly once (v1 cached the last prompt token twice)"
18860        );
18861    }
18862
18863    #[test]
18864    fn generation_is_reproducible_with_seed() {
18865        let run = || {
18866            let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18867            p.generate("hello", 8, None, None).unwrap().token_ids
18868        };
18869        assert_eq!(run(), run());
18870    }
18871
18872    #[test]
18873    fn resetting_sampler_restarts_the_seeded_stream() {
18874        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
18875        let config = SamplerConfig {
18876            seed: Some(1234),
18877            ..SamplerConfig::default()
18878        };
18879        p.set_sampler_config(config.clone());
18880        let first = p.generate("hello", 8, None, None).unwrap().token_ids;
18881        p.set_sampler_config(config);
18882        let second = p.generate("hello", 8, None, None).unwrap().token_ids;
18883        assert_eq!(first, second);
18884    }
18885
18886    #[test]
18887    fn eviction_bounds_the_cache() {
18888        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
18889        p.kv_cache.max_seq_len = 6;
18890        p.sampler_config.temperature = 0.0;
18891        let _ = p.generate("abcd", 12, None, None).unwrap();
18892        assert!(
18893            p.kv_cache.seq_len() <= 6 + 1,
18894            "cache must stay bounded by max_seq_len (got {})",
18895            p.kv_cache.seq_len()
18896        );
18897    }
18898
18899    #[test]
18900    fn confidence_matches_tokens_and_is_a_probability() {
18901        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18902        p.sampler_config.temperature = 0.0;
18903        p.sampler_config.repetition_penalty = 1.0;
18904        let r = p.generate("abcd", 10, None, None).unwrap();
18905        assert_eq!(
18906            r.token_confidence.len(),
18907            r.token_ids.len(),
18908            "one confidence per emitted token"
18909        );
18910        for &c in &r.token_confidence {
18911            assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
18912        }
18913        // top1_prob is a valid softmax probability.
18914        let logits = [1.0f32, 3.0, 0.5, 3.0];
18915        let p0 = top1_prob_t(&logits, 1, 1.0);
18916        let p1 = top1_prob_t(&logits, 3, 1.0);
18917        assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
18918        assert!(p0 > 0.0 && p0 < 1.0);
18919        // Calibration temperature > 1 softens an over-confident peak.
18920        let sharp = top1_prob_t(&logits, 1, 1.0);
18921        let soft = top1_prob_t(&logits, 1, 2.0);
18922        assert!(soft < sharp, "higher temperature lowers peak confidence");
18923    }
18924
18925    #[test]
18926    fn trace_is_opt_in_and_parallels_the_output() {
18927        // Off by default: the runtime is silent unless observation asked.
18928        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18929        p.sampler_config.temperature = 0.0;
18930        p.sampler_config.repetition_penalty = 1.0;
18931        let r = p.generate("abcd", 10, None, None).unwrap();
18932        assert!(r.traces.is_empty(), "trace must be empty unless enabled");
18933
18934        // On: exactly one row per emitted token, aligned with the output.
18935        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18936        p.sampler_config.temperature = 0.0;
18937        p.sampler_config.repetition_penalty = 1.0;
18938        p.set_trace(true);
18939        let r = p.generate("abcd", 10, None, None).unwrap();
18940        assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
18941        for (i, tr) in r.traces.iter().enumerate() {
18942            assert_eq!(tr.t, i, "trace index is sequential");
18943            assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
18944            assert_eq!(
18945                tr.confidence, r.token_confidence[i],
18946                "trace confidence matches the confidence channel"
18947            );
18948            // No dynamic router in this pipeline → no skill, no coherence.
18949            assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
18950        }
18951    }
18952
18953    #[test]
18954    fn explain_prefill_logits_match_greedy_first_token() {
18955        // `cortiq explain` shows the next-token distribution from
18956        // prefill_next_logits; its argmax must equal what greedy generate
18957        // actually emits first — otherwise explain would lie.
18958        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
18959        p.sampler_config.temperature = 0.0;
18960        p.sampler_config.repetition_penalty = 1.0;
18961        let ids = p.tokenizer.encode("abcd");
18962        let logits = p.prefill_next_logits(&ids, None);
18963        let argmax = logits
18964            .iter()
18965            .enumerate()
18966            .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
18967            .unwrap()
18968            .0 as u32;
18969        let r = p.generate("abcd", 1, None, None).unwrap();
18970        assert_eq!(
18971            argmax, r.token_ids[0],
18972            "explain preview must match greedy emit"
18973        );
18974    }
18975
18976    #[test]
18977    fn laguna_shared_expert_is_unconditionally_added() {
18978        let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
18979        let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
18980        let zero_dense = || DenseFfn {
18981            gate_proj: matrix(vec![0.0; 4]),
18982            up_proj: matrix(vec![0.0; 4]),
18983            down_proj: matrix(vec![0.0; 4]),
18984            act: Act::Silu,
18985            down_t: None,
18986            segs: Vec::new(),
18987        };
18988        let shared = DenseFfn {
18989            gate_proj: identity(),
18990            up_proj: identity(),
18991            down_proj: identity(),
18992            act: Act::Silu,
18993            down_t: None,
18994            segs: Vec::new(),
18995        };
18996        let x = [1.0, 2.0];
18997        let expected = dense_ffn(&shared, &x, None);
18998        let moe = MoeFfn {
18999            router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
19000            experts: vec![zero_dense()],
19001            top_k: 1,
19002            norm_topk_prob: true,
19003            router_sigmoid: true,
19004            expert_bias: None,
19005            routed_scaling: 1.0,
19006            route_tau: None,
19007            shared: Some((shared, None)),
19008            stats: std::cell::RefCell::new(Vec::new()),
19009            act_sq: std::cell::RefCell::new(Vec::new()),
19010            act_rows: std::cell::RefCell::new(Vec::new()),
19011            mask: None,
19012            per_expert_scale: None,
19013            router_input_norm: false,
19014            resonance: None,
19015            grown: Vec::new(),
19016        };
19017        let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
19018        for (actual, expected) in actual.iter().zip(expected) {
19019            assert!((actual - expected).abs() < 1e-6);
19020        }
19021    }
19022
19023    /// A tiny MiMo-V2-shaped stack (the M3 fixture): layers [full, sliding,
19024    /// sliding, full]; 4 Q heads over 1 (full) / 2 (sliding) KV heads;
19025    /// head_dim 8 with 4-wide V heads; partial rotary 4 at θ 1e7 (full) /
19026    /// 1e4 (sliding); window 3; learned sinks on the sliding layers; layer
19027    /// 0 a dense FFN, layers 1..3 sigmoid-routed MoE with a selection bias
19028    /// (4 experts, top-2, renormalized, no shared expert). Geometry and
19029    /// sinks go through the same `set_attn_geometry` / `set_layer_sinks`
19030    /// the loader calls.
19031    fn mimo_test_pipeline() -> Pipeline {
19032        let (hs, inter, nh, hd, vd, vocab) = (16usize, 24usize, 4usize, 8usize, 4usize, 64usize);
19033        let kvh = [1usize, 2, 2, 1];
19034        let synth = |n: usize, salt: usize| -> Vec<f32> {
19035            (0..n)
19036                .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
19037                .collect()
19038        };
19039        let qt = |rows: usize, cols: usize, salt: usize| {
19040            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19041        };
19042        let dense = |inter: usize, salt: usize| DenseFfn {
19043            gate_proj: qt(inter, hs, salt),
19044            up_proj: qt(inter, hs, salt + 1),
19045            down_proj: qt(hs, inter, salt + 2),
19046            act: Act::Silu,
19047            down_t: None,
19048            segs: Vec::new(),
19049        };
19050        let layers: Vec<LayerWeights> = (0..4)
19051            .map(|li| LayerWeights {
19052                input_norm: vec![1.0; hs],
19053                post_norm: vec![1.0; hs],
19054                attn_out_norm: None,
19055                ffn_out_norm: None,
19056                layer_scale: None,
19057                ffn: if li == 0 {
19058                    FfnKind::Dense(dense(inter, 50))
19059                } else {
19060                    FfnKind::Moe(MoeFfn {
19061                        router: qt(4, hs, 60 + li),
19062                        experts: (0..4).map(|e| dense(8, 70 + li * 10 + e * 3)).collect(),
19063                        top_k: 2,
19064                        norm_topk_prob: true,
19065                        router_sigmoid: true,
19066                        expert_bias: Some(vec![0.02, -0.03, 0.01, 0.0]),
19067                        routed_scaling: 1.0,
19068                        route_tau: None,
19069                        shared: None,
19070                        stats: std::cell::RefCell::new(Vec::new()),
19071                        act_sq: std::cell::RefCell::new(Vec::new()),
19072                        act_rows: std::cell::RefCell::new(Vec::new()),
19073                        mask: None,
19074                        per_expert_scale: None,
19075                        router_input_norm: false,
19076                        resonance: None,
19077                        grown: Vec::new(),
19078                    })
19079                },
19080                attn: AttnKind::Full {
19081                    wq: qt(nh * hd, hs, li * 10 + 1),
19082                    wk: qt(kvh[li] * hd, hs, li * 10 + 2),
19083                    wv: qt(kvh[li] * vd, hs, li * 10 + 3),
19084                    wo: qt(hs, nh * vd, li * 10 + 4),
19085                    q_norm: None,
19086                    k_norm: None,
19087                    output_gate: false,
19088                    softplus_gate: None,
19089                    bias: None,
19090                },
19091            })
19092            .collect();
19093        let mut p = Pipeline::new(
19094            Tokenizer::byte_level(),
19095            PipelineWeights {
19096                embed_tokens: qt(vocab, hs, 100),
19097                layers,
19098                lm_head: qt(vocab, hs, 200),
19099                final_norm: vec![1.0; hs],
19100            },
19101            hs,
19102            inter,
19103            nh,
19104            1, // header num_kv_heads (the full layers')
19105            hd,
19106            4,
19107            4,
19108            false,
19109            vocab,
19110            1e-6,
19111            1e7,
19112            NormStyle::Qwen,
19113            4096,
19114            SamplerConfig {
19115                seed: Some(7),
19116                ..Default::default()
19117            },
19118        );
19119        // Diagnostics stay off whatever the test environment exports.
19120        p.layer_dump = None;
19121        p.set_rotary(4, 1e7);
19122        p.sliding_layers = Some(vec![false, true, true, false]);
19123        p.swa = Some((3, usize::MAX));
19124        p.rotary_dim_local = Some(4);
19125        p.inv_freq_local = Some(std::sync::Arc::new(attention::rope_inv_freq(4, 1e4)));
19126        p.set_attn_geometry(Some(kvh.to_vec()), Some(vd)).unwrap();
19127        p.set_layer_sinks(1, vec![0.5, -1.0, 1.5, 0.0]).unwrap();
19128        p.set_layer_sinks(2, vec![-0.25, 2.0, 0.75, -1.5]).unwrap();
19129        p
19130    }
19131
19132    #[test]
19133    fn mimo_embedded_prompt_uses_rows_and_never_reuses_token_only_kv() {
19134        let mut p = mimo_test_pipeline();
19135        p.speculative = false;
19136        p.ignore_eos = true;
19137        p.sampler_config.temperature = 0.0;
19138        p.sampler_config.repetition_penalty = 1.0;
19139        let a = vec![3, 5, 7, 9, 11, 13];
19140        let b = vec![4, 8, 12, 16, 20, 24];
19141        let rows: Vec<_> = b.iter().flat_map(|&id| p.embed_id(id)).collect();
19142        let expected = p.generate_from_ids(&b, 8, None, None).unwrap().token_ids;
19143        // Same placeholder IDs as an earlier request are not a cache key
19144        // for different media. The actual rows, not a re-embedding of a,
19145        // must determine the continuation.
19146        let actual = p.generate_from_embeds(&a, &rows, 8, None, None).unwrap().token_ids;
19147        assert_eq!(actual, expected);
19148        assert!(p.kv_history.is_empty());
19149        let mut extended = a.clone();
19150        extended.push(17);
19151        let after_media = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19152        p.reset_session();
19153        let fresh = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19154        assert_eq!(after_media, fresh);
19155        assert!(p.generate_from_embeds(&a, &rows[..rows.len()-1], 1, None, None).is_err());
19156        // Force a real token-prefix reuse opportunity into the media call.
19157        // Those labels are unchanged, but their embeddings now describe a
19158        // different source sequence and every KV row must be rebuilt.
19159        p.reset_session();
19160        p.generate_from_ids(&a, 1, None, None).unwrap();
19161        let mut media_ids = p.kv_history.clone();
19162        assert!(!media_ids.is_empty());
19163        media_ids.extend_from_slice(&[19, 21, 23]);
19164        let source_ids: Vec<_> = (0..media_ids.len()).map(|i| b[i % b.len()]).collect();
19165        let source_rows: Vec<_> = source_ids.iter().flat_map(|&id| p.embed_id(id)).collect();
19166        let mut oracle = mimo_test_pipeline();
19167        oracle.speculative = false;
19168        oracle.ignore_eos = true;
19169        oracle.sampler_config.temperature = 0.0;
19170        oracle.sampler_config.repetition_penalty = 1.0;
19171        let expected = oracle.generate_from_ids(&source_ids, 8, None, None).unwrap().token_ids;
19172        assert_eq!(p.generate_from_embeds(&media_ids, &source_rows, 8, None, None).unwrap().token_ids, expected);
19173        assert!(p.kv_history.is_empty());
19174        let mut bad = rows;
19175        bad[0] = f32::NAN;
19176        assert!(p.generate_from_embeds(&a, &bad, 1, None, None).is_err());
19177    }
19178
19179    fn f32_bits(v: &[f32]) -> Vec<u32> {
19180        v.iter().map(|x| x.to_bits()).collect()
19181    }
19182
19183    /// M3 acceptance: on the MiMo-shaped stack the decode walk (one
19184    /// position at a time through `forward_layers`) and the batched
19185    /// prefill (`prefill_batch_span`, whole prompt and split in two
19186    /// chunks) give bit-identical logits at all 12 positions — per-layer
19187    /// KV heads, narrow V, sinks, the window and the biased sigmoid MoE all
19188    /// agree across the two walks.
19189    #[test]
19190    fn mimo_shaped_decode_matches_prefill_batch_bitwise() {
19191        let mut p = mimo_test_pipeline();
19192        let kv: Vec<usize> = p.kv_cache.layers.iter().map(|l| l.num_kv_heads).collect();
19193        assert_eq!(kv, vec![1, 2, 2, 1]);
19194        assert!(p.kv_cache.layers[0].sinks.is_none() && p.kv_cache.layers[3].sinks.is_none());
19195        assert!(p.kv_cache.layers[1].sinks.is_some() && p.kv_cache.layers[2].sinks.is_some());
19196        let ids: Vec<u32> = (0..12u32).map(|i| (i * 7 + 3) % 64).collect();
19197        let hs = p.hidden_size;
19198        let mut decode = Vec::new();
19199        for (pos, &id) in ids.iter().enumerate() {
19200            let e = p.embed_single(id);
19201            let h = p.forward_layers(&e, pos, None);
19202            decode.push(p.logits_from_hidden(&h));
19203        }
19204        for l in &p.kv_cache.layers {
19205            assert_eq!(l.seq_len, 12);
19206            // V rows are padded to head_dim inside the cache.
19207            assert_eq!(l.head_values(0).len(), 12 * 8);
19208        }
19209        assert!(decode.iter().flatten().all(|v| v.is_finite()));
19210
19211        p.clear_sequence_state();
19212        let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19213        for pos in 0..ids.len() {
19214            let lg = p.logits_from_hidden(&hb[pos * hs..(pos + 1) * hs]);
19215            assert_eq!(
19216                f32_bits(&decode[pos]),
19217                f32_bits(&lg),
19218                "whole prompt, pos {pos}"
19219            );
19220        }
19221
19222        p.clear_sequence_state();
19223        let a = p.prefill_batch_span(PrefillIn::Ids(&ids[..5]), 0, None, 0, p.num_layers);
19224        let b = p.prefill_batch_span(PrefillIn::Ids(&ids[5..]), 5, None, 0, p.num_layers);
19225        for pos in 0..ids.len() {
19226            let row = if pos < 5 {
19227                &a[pos * hs..(pos + 1) * hs]
19228            } else {
19229                &b[(pos - 5) * hs..(pos - 4) * hs]
19230            };
19231            let lg = p.logits_from_hidden(row);
19232            assert_eq!(
19233                f32_bits(&decode[pos]),
19234                f32_bits(&lg),
19235                "two chunks, pos {pos}"
19236            );
19237        }
19238
19239        // The fixture is not degenerate: the sinks and the window each
19240        // change the answer.
19241        let last = |p: &mut Pipeline| {
19242            p.clear_sequence_state();
19243            let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19244            p.logits_from_hidden(&hb[11 * hs..12 * hs])
19245        };
19246        let base = last(&mut p);
19247        let mut no_sinks = mimo_test_pipeline();
19248        for l in &mut no_sinks.kv_cache.layers {
19249            l.sinks = None;
19250        }
19251        assert_ne!(
19252            f32_bits(&last(&mut no_sinks)),
19253            f32_bits(&base),
19254            "sinks are live"
19255        );
19256        let mut wide = mimo_test_pipeline();
19257        wide.swa = Some((64, usize::MAX));
19258        assert_ne!(
19259            f32_bits(&last(&mut wide)),
19260            f32_bits(&base),
19261            "window is live"
19262        );
19263
19264        // Generation runs end to end on the same stack.
19265        p.clear_sequence_state();
19266        p.ignore_eos = true;
19267        let r = p.generate_from_ids(&ids, 4, None, None).unwrap();
19268        assert_eq!(r.token_ids.len(), 4);
19269    }
19270
19271    /// A synthetic MiMo draft stack of `n` layers for `mimo_test_pipeline`
19272    /// (the SWA geometry of its sliding layers: 2 KV heads, head 8 / V 4).
19273    fn mimo_test_mtp(n: usize, gain: f32) -> mimo_mtp::MimoMtp {
19274        let (hs, inter, nh, hd, vd, nkv) = (16usize, 24usize, 4usize, 8usize, 4usize, 2usize);
19275        let synth = |len: usize, salt: usize| -> Vec<f32> {
19276            (0..len)
19277                .map(|i| (((i * 37 + salt * 13 + 3) % 89) as f32 / 89.0 - 0.5) * 0.6 * gain)
19278                .collect()
19279        };
19280        let qt = |rows: usize, cols: usize, salt: usize| {
19281            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19282        };
19283        let layers = (0..n)
19284            .map(|k| {
19285                let s = 500 + k * 40;
19286                let mut kv = crate::kv_cache::LayerKvCache::new(nkv, hd);
19287                kv.sinks = Some(vec![0.3, -0.7, 1.1, 0.0]);
19288                MtpModule {
19289                    enorm: vec![1.0; hs],
19290                    hnorm: vec![1.0; hs],
19291                    eh_proj: qt(hs, 2 * hs, s),
19292                    layer: LayerWeights {
19293                        input_norm: vec![1.0; hs],
19294                        post_norm: vec![1.0; hs],
19295                        attn_out_norm: None,
19296                        ffn_out_norm: None,
19297                        layer_scale: None,
19298                        attn: AttnKind::Full {
19299                            wq: qt(nh * hd, hs, s + 1),
19300                            wk: qt(nkv * hd, hs, s + 2),
19301                            wv: qt(nkv * vd, hs, s + 3),
19302                            wo: qt(hs, nh * vd, s + 4),
19303                            q_norm: None,
19304                            k_norm: None,
19305                            output_gate: false,
19306                            softplus_gate: None,
19307                            bias: None,
19308                        },
19309                        ffn: FfnKind::Dense(DenseFfn {
19310                            gate_proj: qt(inter, hs, s + 5),
19311                            up_proj: qt(inter, hs, s + 6),
19312                            down_proj: qt(hs, inter, s + 7),
19313                            act: Act::Silu,
19314                            down_t: None,
19315                            segs: Vec::new(),
19316                        }),
19317                    },
19318                    final_norm: vec![1.0; hs],
19319                    kv,
19320                }
19321            })
19322            .collect();
19323        mimo_mtp::MimoMtp::from_layers(layers)
19324    }
19325
19326    fn mimo_greedy(p: &mut Pipeline, ids: &[u32], n: usize, spec: bool) -> GenerateResult {
19327        p.clear_sequence_state();
19328        p.speculative = spec;
19329        p.ignore_eos = true;
19330        p.sampler_config.temperature = 0.0;
19331        p.generate_from_ids(ids, n, None, None).unwrap()
19332    }
19333
19334    /// The draft stack's incremental rounds (a few rows per layer, last
19335    /// round's provisional rows dropped) give exactly the teacher-forced
19336    /// table of one causal pass per layer over the whole sequence — the
19337    /// table `tools/mimo_ref.py mtp` computes for variant A: layer k, row
19338    /// j reads (x[j+k+1], norm(h_j)) at RoPE position j.
19339    #[test]
19340    fn mimo_mtp_incremental_rounds_equal_the_teacher_forced_table() {
19341        // Both readings of the backbone hidden: pre-final-norm (default)
19342        // and post-final-norm (`CMF_MIMO_MTP_HIDDEN=post`).
19343        for post in [false, true] {
19344            let mut p = mimo_test_pipeline();
19345            // A non-trivial final norm, so the two readings differ.
19346            p.weights.final_norm = (0..p.hidden_size).map(|i| 0.5 + 0.1 * i as f32).collect();
19347            let mut st0 = mimo_test_mtp(3, 1.0);
19348            st0.post_norm_hidden = post;
19349            p.mimo_mtp = Some(st0);
19350            let ids: Vec<u32> = (0..14u32).map(|i| (i * 11 + 5) % 64).collect();
19351            let hs = p.hidden_size;
19352            let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19353            p.mimo_note_rows(&hb, 0);
19354            let mut st = p.mimo_mtp.take().unwrap();
19355            // Incremental: one round per t through the decode path (later
19356            // tokens from `ids`, the probe's teacher forcing).
19357            let k = 3;
19358            let mut inc = Vec::new();
19359            for t in 0..ids.len() - k - 1 {
19360                inc.push(p.mimo_mtp_draft(&mut st, t, &ids, k));
19361            }
19362            // Reference: per layer, ONE batched causal pass over all rows
19363            // with fresh caches.
19364            let s = ids.len();
19365            let mut reference = vec![vec![0u32; k]; s - k - 1];
19366            let mut fresh = mimo_test_mtp(3, 1.0);
19367            for (layer, m) in fresh.layers.iter_mut().enumerate() {
19368                let n = s - layer - 1;
19369                let mut cats = vec![0.0f32; n * 2 * hs];
19370                for j in 0..n {
19371                    let e = p.embed_single(ids[j + layer + 1]);
19372                    let raw = &hb[j * hs..(j + 1) * hs];
19373                    let g = if post {
19374                        inference::rms_norm(raw, &p.weights.final_norm, p.rms_eps, p.norm_style)
19375                    } else {
19376                        raw.to_vec()
19377                    };
19378                    let (ce, ch) = cats[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
19379                    inference::rms_norm_into(&e, &m.enorm, p.rms_eps, p.norm_style, ce);
19380                    inference::rms_norm_into(&g, &m.hnorm, p.rms_eps, p.norm_style, ch);
19381                }
19382                let mut x = vec![0.0f32; n * hs];
19383                m.eh_proj.matmat(&cats, n, &mut x, None);
19384                p.mimo_mtp_block(m, &mut x, n, 0);
19385                for (t, row) in reference.iter_mut().enumerate() {
19386                    let y = inference::rms_norm(
19387                        &x[t * hs..(t + 1) * hs],
19388                        &m.final_norm,
19389                        p.rms_eps,
19390                        p.norm_style,
19391                    );
19392                    row[layer] = sampler::argmax(&p.lm_head_forward(&y));
19393                }
19394            }
19395            assert_eq!(inc, reference, "post_norm_hidden = {post}");
19396            // Not a degenerate table: the drafts vary.
19397            let distinct: std::collections::HashSet<u32> =
19398                inc.iter().flatten().copied().collect();
19399            assert!(distinct.len() > 3, "{inc:?}");
19400            // Each layer's cache ends holding rows up to the last round start.
19401            let last_t = ids.len() - k - 2;
19402            for m in &st.layers {
19403                assert_eq!(m.kv.seq_len, last_t + 1);
19404            }
19405        }
19406    }
19407
19408    /// Greedy with the MiMo draft stack is the plain greedy stream, token
19409    /// for token — with the real draft layers (low acceptance) and with a
19410    /// drafter that is right most of the time (exercises accepted prefixes
19411    /// of every length, the KV truncation of the rejected rows and the
19412    /// logits hand-off to the loop top), under the default repetition
19413    /// penalty.
19414    #[test]
19415    fn mimo_speculative_greedy_equals_plain_greedy() {
19416        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19417        let ids: Vec<u32> = (0..9u32).map(|i| (i * 7 + 3) % 64).collect();
19418        let n = 24;
19419        let mut p = mimo_test_pipeline();
19420        let plain = mimo_greedy(&mut p, &ids, n, false);
19421        assert_eq!(plain.mtp_drafted, 0);
19422        assert_eq!(plain.token_ids.len(), n);
19423        let plain_kv = p.kv_cache.layers[0].seq_len;
19424
19425        // Real draft layers.
19426        p.mimo_mtp = Some(mimo_test_mtp(3, 1.0));
19427        let spec = mimo_greedy(&mut p, &ids, n, true);
19428        assert!(spec.mtp_drafted > 0, "the round must draft");
19429        assert_eq!(spec.token_ids, plain.token_ids);
19430        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19431
19432        // A drafter reading the true continuation with every fifth token
19433        // wrong: accepted prefixes of 0..=3 all occur.
19434        let mut truth: Vec<u32> = ids.clone();
19435        truth.extend(&plain.token_ids);
19436        let mut noisy = truth.clone();
19437        for (i, t) in noisy.iter_mut().enumerate() {
19438            if i % 5 == 0 {
19439                *t = (*t + 1) % 64;
19440            }
19441        }
19442        let mut st = mimo_test_mtp(3, 1.0);
19443        st.draft_override = Some(noisy);
19444        p.mimo_mtp = Some(st);
19445        let spec = mimo_greedy(&mut p, &ids, n, true);
19446        assert_eq!(spec.token_ids, plain.token_ids);
19447        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19448        let stats = p.mimo_mtp.as_ref().unwrap().stats.clone();
19449        assert_eq!(stats.accepted as usize, spec.mtp_accepted);
19450        assert!(spec.mtp_accepted > 0 && spec.mtp_accepted < spec.mtp_drafted);
19451        assert!(
19452            stats.accept_hist.iter().filter(|&&c| c > 0).count() >= 3,
19453            "{:?}",
19454            stats.accept_hist
19455        );
19456        assert!(stats.tokens_per_round() > 1.5, "{}", stats.line());
19457
19458        // A perfect drafter: every draft accepted, rounds of K+1 tokens,
19459        // and the budget is never overrun.
19460        let mut st = mimo_test_mtp(3, 1.0);
19461        st.draft_override = Some(truth);
19462        p.mimo_mtp = Some(st);
19463        let spec = mimo_greedy(&mut p, &ids, n, true);
19464        assert_eq!(spec.token_ids, plain.token_ids);
19465        assert_eq!(spec.mtp_accepted, spec.mtp_drafted);
19466        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
19467
19468        // CMF_MTP=0 path: the stack is attached but idle.
19469        let off = mimo_greedy(&mut p, &ids, n, false);
19470        assert_eq!(off.token_ids, plain.token_ids);
19471        assert_eq!(off.mtp_drafted, 0);
19472    }
19473
19474    /// The wgpu graphs carry MiMo-V2's attention per layer (KV heads,
19475    /// narrow V, sinks, windows, two RoPE tables): no attention-level
19476    /// decline for it any more, and the geometry each layer hands the
19477    /// graph is exactly what the CPU attention reads for that layer. The
19478    /// descriptive reasons stay (the Metal graphs and the q1 dropin still
19479    /// decline on them), and what the per-layer geometry cannot express
19480    /// keeps a named wgpu decline.
19481
19482    #[test]
19483    fn mimo_shaped_model_rides_the_wgpu_graph_geometry() {
19484        let p = mimo_test_pipeline();
19485        assert_eq!(
19486            p.graph_attn_decline_reason(),
19487            Some("per-layer KV head counts")
19488        );
19489        assert_eq!(p.wgpu_graph_attn_decline(), None);
19490        let g0 = p.graph_attn_geom(0).expect("full layer geometry");
19491        assert_eq!(
19492            (g0.nkv, g0.dv, g0.rd, g0.window, g0.sink.is_some()),
19493            (1, 4, 4, None, false)
19494        );
19495        assert_eq!(g0.invf, p.inv_freq.as_slice());
19496        let g1 = p.graph_attn_geom(1).expect("sliding layer geometry");
19497        assert_eq!((g1.nkv, g1.dv, g1.rd, g1.window), (2, 4, 4, Some(3)));
19498        assert_eq!(g1.sink, Some(&[0.5f32, -1.0, 1.5, 0.0][..]));
19499        assert_eq!(g1.invf, p.inv_freq_local.as_ref().unwrap().as_slice());
19500        assert_ne!(g0.invf, g1.invf, "two RoPE tables");
19501        let g3 = p.graph_attn_geom(3).expect("full layer geometry");
19502        assert_eq!((g3.nkv, g3.window, g3.sink.is_some()), (1, None, false));
19503
19504        // No wgpu device in this process: the builders run and decline on
19505        // the (f32, unmapped) experts — never with an attention line.
19506        let emb = p.embed_single(3);
19507        let mut lg = Vec::new();
19508        assert!(
19509            p.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, p.num_layers)
19510                .is_none()
19511        );
19512        let mut hid = emb.clone();
19513        assert_eq!(
19514            p.try_batch_graph_wgpu(&mut hid, &[0], 1, None),
19515            crate::gpu::BatchGraphOutcome::Declined
19516        );
19517        assert_eq!(hid, emb, "a declined batch graph leaves the rows untouched");
19518        assert!(p.try_multi_burst(3, 0, 4).is_none());
19519        assert!(
19520            p.graph_declines().is_empty(),
19521            "no attention decline logged: {:?}",
19522            p.graph_declines()
19523        );
19524        // (No assertion on graph_prefill_preferred: with no attention
19525        // decline it follows the device — a test process that brought a
19526        // wgpu adapter up routes this resident MoE through the graph.)
19527
19528        let plain = || create_test_pipeline(8, 16, 2, 1, 4, 2, 32);
19529        assert_eq!(plain().graph_attn_decline_reason(), None);
19530        assert_eq!(plain().wgpu_graph_attn_decline(), None);
19531        assert!(
19532            plain().graph_attn_geom(0).is_none(),
19533            "uniform models keep the historical arms"
19534        );
19535        let mut q = plain();
19536        q.set_layer_sinks(1, vec![0.25, -0.25]).unwrap();
19537        assert_eq!(
19538            q.graph_attn_decline_reason(),
19539            Some("learned attention sinks")
19540        );
19541        assert_eq!(
19542            q.graph_attn_geom(1).unwrap().sink,
19543            Some(&[0.25f32, -0.25][..])
19544        );
19545        let mut q = plain();
19546        q.set_attn_geometry(None, Some(2)).unwrap();
19547        assert_eq!(
19548            q.graph_attn_decline_reason(),
19549            Some("V heads narrower than Q/K heads")
19550        );
19551        assert_eq!(q.graph_attn_geom(0).unwrap().dv, 2);
19552        let mut q = plain();
19553        q.sliding_layers = Some(vec![true, false]);
19554        q.swa = Some((4, usize::MAX));
19555        assert_eq!(q.graph_attn_decline_reason(), Some("sliding-window layers"));
19556        assert_eq!(q.graph_attn_geom(0).unwrap().window, Some(4));
19557        assert_eq!(q.graph_attn_geom(1).unwrap().window, None);
19558
19559        // Outside the per-layer geometry: a named wgpu decline, logged
19560        // once per site.
19561        let mut q = mimo_test_pipeline();
19562        q.rope_scale = 2.0;
19563        assert_eq!(
19564            q.wgpu_graph_attn_decline(),
19565            Some("scaled RoPE positions with per-layer geometry")
19566        );
19567        let emb = q.embed_single(3);
19568        assert!(
19569            q.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, q.num_layers)
19570                .is_none()
19571        );
19572        let _ = q.try_token_graph_wgpu_steps(&emb, 1, &mut lg, 1, None, None, 0, q.num_layers);
19573        let lines = q.graph_declines();
19574        assert_eq!(
19575            lines
19576                .iter()
19577                .filter(|l| l.starts_with("wgpu token graph") && l.contains("scaled RoPE"))
19578                .count(),
19579            1,
19580            "{lines:?}"
19581        );
19582    }
19583
19584    #[test]
19585    fn mimo_verify_rewind_preserves_lagging_host_caches() {
19586        let mut p = mimo_test_pipeline();
19587        for (li, layer) in p.kv_cache.layers.iter_mut().enumerate() {
19588            let row = vec![0.0; layer.num_kv_heads * layer.head_dim];
19589            for _ in 0..if li == 0 { 2 } else { 12 } {
19590                layer.append(&row, &row, &[]);
19591            }
19592        }
19593        p.mimo_verify_rewind(9).unwrap();
19594        assert_eq!(p.kv_cache.layers[0].seq_len, 2);
19595        for layer in &p.kv_cache.layers[1..] {
19596            assert_eq!(layer.seq_len, 9);
19597        }
19598    }
19599
19600    /// CMF_LAYER_DUMP: the decode walk and the batched prefill both write
19601    /// every (position, layer) hidden, the two sets agree byte for byte,
19602    /// and the last layer's file is the stack output.
19603    #[test]
19604    fn layer_dump_covers_every_position_and_layer_on_both_walks() {
19605        let dir = std::env::temp_dir().join(format!("cmf-layer-dump-{}", std::process::id()));
19606        let _ = std::fs::remove_dir_all(&dir);
19607        let mut p = mimo_test_pipeline();
19608        let hs = p.hidden_size;
19609        let ids = [5u32, 9, 11, 2, 40];
19610        p.layer_dump = Some(dir.join("decode"));
19611        for (pos, &id) in ids.iter().enumerate() {
19612            let e = p.embed_single(id);
19613            let _ = p.forward_layers(&e, pos, None);
19614        }
19615        p.clear_sequence_state();
19616        p.layer_dump = Some(dir.join("prefill"));
19617        let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19618        for pos in 0..ids.len() {
19619            for li in 0..p.num_layers {
19620                let name = format!("p{pos:06}_l{li:02}.f32");
19621                let a = std::fs::read(dir.join("decode").join(&name)).unwrap();
19622                let b = std::fs::read(dir.join("prefill").join(&name)).unwrap();
19623                assert_eq!(a.len(), hs * 4, "{name}");
19624                assert_eq!(a, b, "{name}");
19625            }
19626        }
19627        let last = std::fs::read(dir.join("prefill").join("p000004_l03.f32")).unwrap();
19628        let vals: Vec<f32> = last
19629            .chunks(4)
19630            .map(|c| f32::from_le_bytes(c.try_into().unwrap()))
19631            .collect();
19632        assert_eq!(f32_bits(&vals), f32_bits(&hb[4 * hs..5 * hs]));
19633        let _ = std::fs::remove_dir_all(&dir);
19634    }
19635
19636    #[test]
19637    fn attn_geometry_and_sinks_are_validated() {
19638        let mut p = create_test_pipeline(8, 16, 4, 2, 4, 2, 32);
19639        assert!(
19640            p.set_attn_geometry(Some(vec![2]), None).is_err(),
19641            "one entry per layer"
19642        );
19643        assert!(
19644            p.set_attn_geometry(Some(vec![2, 3]), None).is_err(),
19645            "3 does not divide 4"
19646        );
19647        assert!(p.set_attn_geometry(Some(vec![2, 0]), None).is_err());
19648        assert!(p.set_attn_geometry(None, Some(0)).is_err());
19649        assert!(
19650            p.set_attn_geometry(None, Some(5)).is_err(),
19651            "V wider than the head"
19652        );
19653        p.set_attn_geometry(None, Some(4)).unwrap();
19654        assert_eq!(
19655            p.v_head_dim, None,
19656            "v_head_dim == head_dim is the uniform case"
19657        );
19658        p.set_layer_sinks(1, vec![0.1; 4]).unwrap();
19659        p.set_attn_geometry(Some(vec![1, 4]), None).unwrap();
19660        assert_eq!(p.kv_cache.layers[0].num_kv_heads, 1);
19661        assert_eq!(p.kv_cache.layers[1].num_kv_heads, 4);
19662        assert!(
19663            p.kv_cache.layers[1].sinks.is_some(),
19664            "a reshape keeps the layer's sinks"
19665        );
19666        assert_eq!(p.layer_geom(1).0, 4);
19667        assert!(
19668            p.set_layer_sinks(0, vec![0.0; 3]).is_err(),
19669            "one sink per Q head"
19670        );
19671        assert!(p.set_layer_sinks(7, vec![0.0; 4]).is_err());
19672        assert!(p.set_layer_sinks(0, vec![f32::NAN, 0.0, 0.0, 0.0]).is_err());
19673    }
19674
19675    /// The O(1) Nyström state replaces a plain full-context softmax; it
19676    /// must never be armed on a sliding, sink or narrow-V layer.
19677    #[test]
19678    fn o1_is_never_armed_on_sink_window_or_narrow_v_layers() {
19679        let cfg = || {
19680            Some(crate::nystrom::O1Cfg {
19681                layers: crate::nystrom::O1Layers::All,
19682                m: 4,
19683                w: 8,
19684                sink: 2,
19685                rect: crate::nystrom::O1Rect::Aggregate,
19686            })
19687        };
19688        let mut p = mimo_test_pipeline();
19689        p.set_o1(cfg());
19690        assert!(!p.o1_active(), "every MiMo-shaped layer is ineligible");
19691        let mut q = create_test_pipeline(8, 16, 2, 1, 4, 3, 64);
19692        q.set_layer_sinks(1, vec![0.0, 0.0]).unwrap();
19693        q.sliding_layers = Some(vec![false, false, true]);
19694        q.swa = Some((4, usize::MAX));
19695        q.set_o1(cfg());
19696        assert_eq!(q.o1_flags, vec![true, false, false]);
19697    }
19698
19699    #[test]
19700    fn o1_batch_transition_publishes_one_epoch_before_serial_handoff() {
19701        const B: usize = 19;
19702        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19703        p.set_o1(Some(crate::nystrom::O1Cfg {
19704            layers: crate::nystrom::O1Layers::All,
19705            m: 4,
19706            w: 8,
19707            sink: 2,
19708            rect: crate::nystrom::O1Rect::Aggregate,
19709        }));
19710        p.o1_begin_with_prefix(Some(B));
19711        let ids: Vec<u32> = (0..B as u32).collect();
19712        let _ = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19713
19714        assert_eq!(p.o1_epoch, 1, "all layers publish one completed transition");
19715        assert!(p.kv_cache.layers.iter().all(|l| l.o1_sealed()));
19716        let next = p.embed_single(B as u32);
19717        let _ = p.forward_layers(&next, B, None);
19718        assert_eq!(p.o1_epoch, 1, "sealed handoff must not republish the epoch");
19719    }
19720
19721    #[test]
19722    fn o1_pair_transition_commits_scratch_before_epoch_publication() {
19723        const B: usize = 19;
19724        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19725        // Keep a real recurrent layer ahead of the Full O(1) layer so the
19726        // pair test observes the GDN lane-2 scratch swap at the same
19727        // boundary, rather than only exercising an artificial scratch vec.
19728        let gdn_cfg = crate::linear_core::GdnCfg {
19729            num_v_heads: 2,
19730            num_k_heads: 1,
19731            key_head_dim: 2,
19732            value_head_dim: 4,
19733            conv_kernel: 3,
19734            hidden_size: 8,
19735            rms_eps: 1e-6,
19736            output_gate_sigmoid: false,
19737        };
19738        let synth = |n: usize, salt: usize| -> Vec<f32> {
19739            (0..n)
19740                .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
19741                .collect()
19742        };
19743        let qt = |rows: usize, cols: usize, salt: usize| {
19744            crate::qtensor::QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19745        };
19746        let c_dim = gdn_cfg.conv_dim();
19747        let vd = gdn_cfg.num_v_heads * gdn_cfg.value_head_dim;
19748        p.weights.layers[0].attn = AttnKind::LinearGdn(crate::linear_core::GdnWeights {
19749            in_proj_qkv: qt(c_dim, 8, 1),
19750            in_proj_z: qt(vd, 8, 2),
19751            in_proj_a: qt(gdn_cfg.num_v_heads, 8, 3),
19752            in_proj_b: qt(gdn_cfg.num_v_heads, 8, 4),
19753            conv1d: synth(c_dim * gdn_cfg.conv_kernel, 5),
19754            a_log: vec![0.2, 0.5],
19755            dt_bias: synth(gdn_cfg.num_v_heads, 6),
19756            norm: vec![1.0; gdn_cfg.value_head_dim],
19757            out_proj: qt(8, vd, 7),
19758        });
19759        p.gdn_cfg = Some(gdn_cfg);
19760        p.set_o1(Some(crate::nystrom::O1Cfg {
19761            layers: crate::nystrom::O1Layers::All,
19762            m: 4,
19763            w: 8,
19764            sink: 2,
19765            rect: crate::nystrom::O1Rect::Aggregate,
19766        }));
19767        p.o1_begin_with_prefix(Some(B));
19768        for pos in 0..B - 2 {
19769            let emb = p.embed_single(pos as u32);
19770            let _ = p.forward_layers(&emb, pos, None);
19771        }
19772        let lane1_state = p.kv_cache.layers[0].linear_state.clone();
19773
19774        let e1 = p.embed_single((B - 2) as u32);
19775        let e2 = p.embed_single((B - 1) as u32);
19776        let _ = p.forward_pair(&e1, &e2, B - 2);
19777
19778        assert_eq!(p.o1_epoch, 1, "pair crossing B publishes one epoch");
19779        assert!(
19780            p.kv_cache
19781                .layers
19782                .iter()
19783                .enumerate()
19784                .all(|(li, l)| !p.o1_flags[li] || l.o1_sealed())
19785        );
19786        assert!(!p.kv_cache.layers[0].linear_state.is_empty());
19787        assert_ne!(
19788            p.kv_cache.layers[0].linear_state, lane1_state,
19789            "real pair must commit GDN lane 2 before returning"
19790        );
19791        assert!(p.kv_cache.layers[0].linear_scratch.is_empty());
19792        let next = p.embed_single(B as u32);
19793        let _ = p.forward_layers(&next, B, None);
19794        assert_eq!(p.o1_epoch, 1, "serial continuation must reuse the epoch");
19795    }
19796
19797    #[test]
19798    fn o1_error_observation_stays_terminal_until_reset() {
19799        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19800        p.set_o1(Some(crate::nystrom::O1Cfg {
19801            layers: crate::nystrom::O1Layers::All,
19802            m: 4,
19803            w: 8,
19804            sink: 2,
19805            rect: crate::nystrom::O1Rect::Aggregate,
19806        }));
19807        p.o1_begin();
19808        p.kv_cache.layers[0].o1_abort("synthetic transition failure".into());
19809
19810        assert!(p.o1_seal_checked().is_err());
19811        assert!(
19812            p.o1_seal_checked().is_err(),
19813            "retry must see the sticky error"
19814        );
19815        let k = vec![0.2f32; 4];
19816        let v = vec![0.3f32; 4];
19817        p.kv_cache.layers[0].append(&k, &v, &[]);
19818        assert_eq!(p.kv_cache.layers[0].seq_len, 0);
19819
19820        p.reset_session();
19821        p.o1_begin();
19822        p.kv_cache.layers[0].append(&k, &v, &[]);
19823        assert_eq!(p.kv_cache.layers[0].seq_len, 1);
19824    }
19825
19826    #[test]
19827    fn nll_graph_failure_is_terminal_and_request_is_reusable() {
19828        let ids = vec![1u32, 2, 3, 4, 5, 6];
19829        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19830        p.graph_logits = Some(vec![123.0]);
19831        p.graph_want_logits = true;
19832        p.graph_failed
19833            .store(true, std::sync::atomic::Ordering::Relaxed);
19834        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
19835        let err = p.nll_ids_from(&ids, 0).expect_err("prior graph failure");
19836        assert!(err.contains("before NLL"));
19837        assert!(p.graph_logits.is_none());
19838        assert!(!p.graph_want_logits);
19839        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19840        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19841
19842        let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19843        let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
19844        let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
19845        assert_eq!(actual.1, expected.1);
19846        assert!((actual.0 - expected.0).abs() < 1e-9);
19847    }
19848
19849    #[test]
19850    fn nll_forward_failure_discards_partial_score_and_clears_sidechannels() {
19851        let ids = vec![1u32, 2, 3, 4, 5, 6];
19852        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19853        p.nll_test_fail_at = Some(1);
19854        let err = p
19855            .nll_ids_from(&ids, 0)
19856            .expect_err("one-shot forward failure");
19857        assert!(err.contains("forward") || err.contains("score row"));
19858        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19859        assert!(!p.graph_want_logits);
19860        assert!(p.graph_logits.is_none());
19861        assert!(p.kv_history.is_empty());
19862
19863        let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19864        let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
19865        let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
19866        assert_eq!(actual.1, expected.1);
19867        assert!((actual.0 - expected.0).abs() < 1e-9);
19868    }
19869
19870    #[test]
19871    fn nll_serial_failure_before_first_row_is_reported() {
19872        let ids = vec![1u32, 2, 3, 4];
19873        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19874        p.nll_test_force_serial = true;
19875        p.nll_test_fail_at = Some(0);
19876        let err = p.nll_ids_from(&ids, 0).expect_err("serial forward failure");
19877        assert!(err.contains("serial forward"));
19878        assert!(p.kv_history.is_empty());
19879        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19880        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19881    }
19882
19883    #[test]
19884    fn ffn_probe_failure_discards_recorder_and_state() {
19885        let ids = vec![1u32, 2, 3, 4];
19886        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19887        p.nll_test_fail_at = Some(0);
19888        let err = p
19889            .probe_ffn_mass_batch(&ids)
19890            .expect_err("probe forward failure");
19891        assert!(err.contains("NLL"));
19892        assert!(FFN_PROBE.with(|probe| probe.borrow().is_none()));
19893        assert!(p.kv_history.is_empty());
19894        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19895    }
19896
19897    #[test]
19898    fn nll_test_controls_are_pipeline_scoped() {
19899        let ids = vec![1u32, 2, 3, 4];
19900        let mut failing = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19901        let mut unaffected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19902        failing.nll_test_force_serial = true;
19903        failing.nll_test_fail_at = Some(0);
19904
19905        assert!(!failing.can_prefill_batched());
19906        assert!(unaffected.can_prefill_batched());
19907        let expected = unaffected
19908            .nll_ids_from(&ids, 0)
19909            .expect("unaffected pipeline remains usable");
19910        let err = failing
19911            .nll_ids_from(&ids, 0)
19912            .expect_err("failure injection belongs to failing pipeline");
19913        assert!(err.contains("serial forward"));
19914        assert!(failing.nll_test_fail_at.is_none());
19915        assert!(unaffected.can_prefill_batched());
19916        let actual = unaffected
19917            .nll_ids_from(&ids, 0)
19918            .expect("unaffected pipeline remains reusable");
19919        assert_eq!(actual.1, expected.1);
19920        assert!((actual.0 - expected.0).abs() < 1e-9);
19921    }
19922
19923    #[test]
19924    fn forward_ids_failure_channel_is_terminal_and_reusable() {
19925        let ids = vec![1u32, 2, 3, 4, 5, 6];
19926        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19927        p.graph_logits = Some(vec![123.0]);
19928        p.graph_want_logits = true;
19929        p.graph_failed
19930            .store(true, std::sync::atomic::Ordering::Relaxed);
19931        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
19932
19933        let err = p
19934            .forward_ids(&ids, None)
19935            .expect_err("a failed forward must not become a valid head result");
19936        assert!(err.contains("forward_ids setup"));
19937        assert!(p.graph_logits.is_none());
19938        assert!(!p.graph_want_logits);
19939        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
19940        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
19941        assert_eq!(p.kv_cache.seq_len(), 0);
19942
19943        let expected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64)
19944            .forward_ids(&ids, None)
19945            .expect("fresh forward_ids");
19946        let actual = p
19947            .forward_ids(&ids, None)
19948            .expect("pipeline remains reusable after a failed forward");
19949        assert_eq!(actual.len(), expected.len());
19950        assert!(
19951            actual
19952                .iter()
19953                .zip(expected)
19954                .all(|(a, b)| (a - b).abs() < 1e-9)
19955        );
19956        assert_eq!(p.kv_cache.seq_len(), ids.len());
19957    }
19958
19959    #[test]
19960    fn sigmoid_router_floor_is_explicit_per_architecture() {
19961        // GLM-5's noaux_tc reference uses +1e-20 while the generic
19962        // LFM2-compatible path uses +1e-6.  At low (but representable)
19963        // sigmoid scores, silently sharing the latter changes expert weights
19964        // by orders of magnitude and can make a routed layer look coherent
19965        // while discarding its expert contribution.
19966        let zero = || DenseFfn {
19967            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19968            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19969            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
19970            act: Act::Silu,
19971            down_t: None,
19972            segs: Vec::new(),
19973        };
19974        let m = MoeFfn {
19975            router: QTensor::from_f32(vec![0.0; 4], 2, 2),
19976            experts: vec![zero(), zero()],
19977            top_k: 1,
19978            norm_topk_prob: true,
19979            router_sigmoid: true,
19980            expert_bias: None,
19981            routed_scaling: 2.5,
19982            route_tau: None,
19983            shared: None,
19984            stats: std::cell::RefCell::new(Vec::new()),
19985            act_sq: std::cell::RefCell::new(Vec::new()),
19986            act_rows: std::cell::RefCell::new(Vec::new()),
19987            mask: None,
19988            per_expert_scale: None,
19989            router_input_norm: false,
19990            resonance: None,
19991            grown: Vec::new(),
19992        };
19993        let logits = [-20.0f32, -20.0];
19994        let (_, p, glm_wsum) = moe_route_with_eps(&logits, &m, None, 1e-20);
19995        let (_, _, generic_wsum) = moe_route(&logits, &m, None);
19996        let expected = (p[0] + 1e-20) / m.routed_scaling;
19997        assert!((glm_wsum - expected).abs() < 1e-15);
19998        assert!(generic_wsum > glm_wsum * 100.0);
19999    }
20000
20001    #[test]
20002    fn resonance_scores_match_formula_and_stable_tie() {
20003        let r = Resonance {
20004            // Three descriptors, hidden=2, one projection row each.
20005            mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0],
20006            u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
20007            k: 1,
20008            bias: vec![1.5, 0.5, 0.0],
20009            shell: Vec::new(),
20010        };
20011        let x = [1.0f32, 1.0];
20012        let mut got = vec![0.0; 3];
20013        r.scores(&x, &mut got);
20014        // Expert 0 and 1 are an exact score tie; the CPU top-1 contract uses
20015        // the lower index.  The values also check d² - (U·d)², not just tie
20016        // ordering.
20017        assert!((got[0] - 0.5).abs() < 1e-6);
20018        assert!((got[1] - 0.5).abs() < 1e-6);
20019        assert!(got[2].abs() < 1e-6);
20020        let best = got
20021            .iter()
20022            .enumerate()
20023            .max_by(|(ia, a), (ib, b)| a.partial_cmp(b).unwrap().then(ib.cmp(ia)))
20024            .map(|(i, _)| i);
20025        assert_eq!(best, Some(0));
20026        assert!(got.iter().all(|v| v.is_finite()));
20027    }
20028
20029    /// The growth shell (spec §2): a grown expert whose reconstruction
20030    /// error lies outside its shell scores −∞, one inside keeps the exact
20031    /// resonance score, trunk rows (`+inf` shell) are bit-identical to the
20032    /// shell-less computation; the process-wide switch disables it.
20033    #[test]
20034    fn resonance_shell_masks_outside_keeps_inside_and_trunk_bits() {
20035        // hidden = 2, rank 1. Experts 0/1 = trunk (shell +inf); 2 and 3 =
20036        // grown, the same descriptor (μ = (0, 1), u = (1, 1)) with shells
20037        // 6.0 and 0.25. At x' = (3, 0): d = (3, −1), d² = 10, proj =
20038        // (3 − 1)² = 4, err = 6 exactly — on the boundary of expert 2's
20039        // shell (kept: the rule is strict `>`), outside expert 3's.
20040        let plain = Resonance {
20041            mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 1.0],
20042            u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0, 1.0],
20043            k: 1,
20044            bias: vec![1.5, 0.5, 0.0, 0.0],
20045            shell: Vec::new(),
20046        };
20047        let shelled = Resonance {
20048            mu: plain.mu.clone(),
20049            u: plain.u.clone(),
20050            k: 1,
20051            bias: plain.bias.clone(),
20052            shell: vec![f32::INFINITY, f32::INFINITY, 6.0, 0.25],
20053        };
20054        assert!(!plain.has_shell());
20055        assert!(shelled.has_shell());
20056        let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
20057        let (mut a, mut b) = (vec![0.0; 4], vec![0.0; 4]);
20058        set_growth_shell(Some(true));
20059        assert!(growth_shell_enabled());
20060        // x = (1, 1): the grown experts reconstruct it exactly (err 0):
20061        // inside both shells, every row the shell-less bits.
20062        let x = [1.0f32, 1.0];
20063        plain.scores(&x, &mut a);
20064        shelled.scores(&x, &mut b);
20065        assert_eq!(bits(&a), bits(&b), "inside every shell: unchanged");
20066        assert!(a[2] == 0.0 && a[3] == 0.0);
20067        // x' = (3, 0): expert 3 → −∞, expert 2 (err == shell) and the
20068        // trunk rows keep their exact bits.
20069        let xo = [3.0f32, 0.0];
20070        plain.scores(&xo, &mut a);
20071        shelled.scores(&xo, &mut b);
20072        assert_eq!(a[2], -6.0);
20073        assert_eq!(a[3], -6.0);
20074        assert_eq!(bits(&a[..3]), bits(&b[..3]), "trunk rows + the boundary row");
20075        assert_eq!(b[3], f32::NEG_INFINITY, "outside the shell: −∞");
20076        assert_eq!(shelled.effective_shell(4), shelled.shell);
20077        // The switch (`CMF_GROWTH_SHELL=off` / `growth-eval --shell off`):
20078        // all +inf, the shell-less bits everywhere.
20079        set_growth_shell(Some(false));
20080        assert!(!growth_shell_enabled());
20081        assert_eq!(shelled.effective_shell(4), vec![f32::INFINITY; 4]);
20082        shelled.scores(&xo, &mut b);
20083        assert_eq!(bits(&a), bits(&b));
20084        set_growth_shell(None);
20085        // A shell vector shorter than the expert count masks nothing
20086        // beyond it (a legacy layer whose tail has no shell).
20087        let short = Resonance {
20088            shell: vec![f32::INFINITY, f32::INFINITY],
20089            ..shelled
20090        };
20091        set_growth_shell(Some(true));
20092        short.scores(&xo, &mut b);
20093        assert_eq!(bits(&a), bits(&b));
20094        set_growth_shell(None);
20095    }
20096
20097    /// `moe_route` with −∞ logits (a grown expert outside its shell):
20098    /// top-1 is the best finite expert with weight exactly 1.0 on both
20099    /// the softmax and the sigmoid path; all −∞ degrades to uniform.
20100    #[test]
20101    fn moe_route_neg_inf_logits_pick_the_best_finite_expert_with_weight_one() {
20102        let zero = || DenseFfn {
20103            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20104            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20105            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20106            act: Act::Silu,
20107            down_t: None,
20108            segs: Vec::new(),
20109        };
20110        let moe = |sigmoid: bool| MoeFfn {
20111            router: QTensor::from_f32(vec![0.0; 8], 4, 2),
20112            experts: vec![zero(), zero(), zero(), zero()],
20113            top_k: 1,
20114            norm_topk_prob: true,
20115            router_sigmoid: sigmoid,
20116            expert_bias: None,
20117            routed_scaling: 1.0,
20118            route_tau: None,
20119            shared: None,
20120            stats: std::cell::RefCell::new(Vec::new()),
20121            act_sq: std::cell::RefCell::new(Vec::new()),
20122            act_rows: std::cell::RefCell::new(Vec::new()),
20123            mask: None,
20124            per_expert_scale: None,
20125            router_input_norm: false,
20126            resonance: None,
20127            grown: Vec::new(),
20128        };
20129        let logits = [-1.0f32, f32::NEG_INFINITY, -0.5, f32::NEG_INFINITY];
20130        for sigmoid in [false, true] {
20131            let m = moe(sigmoid);
20132            let (idx, p, wsum) = moe_route(&logits, &m, None);
20133            assert_eq!(idx, vec![2], "sigmoid {sigmoid}");
20134            assert_eq!(p[1], 0.0);
20135            assert_eq!(p[3], 0.0);
20136            assert!(p[2] > p[0] && p[0] > 0.0);
20137            assert!(p.iter().all(|v| v.is_finite()));
20138            let w = p[2] / wsum;
20139            if sigmoid {
20140                // The sigmoid renorm keeps its reference floor (+1e-6).
20141                assert!((w - 1.0).abs() < 1e-5, "sigmoid: weight {w}");
20142            } else {
20143                assert_eq!(w, 1.0, "softmax: the renormalized top-1 weight is exactly 1");
20144            }
20145            // Masked experts stay masked even when they are the only ones
20146            // "admitted" by an allow-list that covers everything.
20147            let (idx, _, _) = moe_route(&logits, &m, Some(&[true, true, true, true]));
20148            assert_eq!(idx, vec![2]);
20149        }
20150        // A finite expert always beats −∞ whatever the bias / order.
20151        let m = moe(false);
20152        let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, f32::NEG_INFINITY, f32::NEG_INFINITY, -9.0], &m, None);
20153        assert_eq!(idx, vec![3]);
20154        // Every expert at −∞ (cannot happen on a grown file — trunk rows
20155        // have no shell): uniform, finite, lowest index.
20156        let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 4], &m, None);
20157        assert_eq!(idx, vec![0]);
20158        assert!(p.iter().all(|&v| v == 0.25));
20159        assert!(wsum.is_finite() && wsum > 0.0);
20160    }
20161
20162    /// The resonance router (top-1) selects by the raw score as the
20163    /// trainer and the graph do — not by softmax probabilities, where two
20164    /// scores closer than 2^-25 collapse to the same `exp(l − max) = 1.0`
20165    /// and the LOWER index wins a token whose score is strictly smaller.
20166    #[test]
20167    fn resonance_route_picks_the_raw_argmax_not_the_softmax_tie() {
20168        let zero = || DenseFfn {
20169            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20170            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20171            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20172            act: Act::Silu,
20173            down_t: None,
20174            segs: Vec::new(),
20175        };
20176        let moe = |resonant: bool, norm_topk: bool| MoeFfn {
20177            router: QTensor::from_f32(vec![0.0; 6], 3, 2),
20178            experts: vec![zero(), zero(), zero()],
20179            top_k: 1,
20180            norm_topk_prob: norm_topk,
20181            router_sigmoid: false,
20182            expert_bias: None,
20183            routed_scaling: 1.0,
20184            route_tau: None,
20185            shared: None,
20186            stats: std::cell::RefCell::new(Vec::new()),
20187            act_sq: std::cell::RefCell::new(Vec::new()),
20188            act_rows: std::cell::RefCell::new(Vec::new()),
20189            mask: None,
20190            per_expert_scale: None,
20191            router_input_norm: false,
20192            resonance: resonant.then(|| Resonance {
20193                mu: vec![0.0; 6],
20194                u: Vec::new(),
20195                k: 0,
20196                bias: vec![0.0; 3],
20197                shell: Vec::new(),
20198            }),
20199            grown: Vec::new(),
20200        };
20201        // lo = −0.1, hi = the next f32 towards zero: hi − lo = 2^-27 <
20202        // 2^-25, so exp(lo − hi) rounds to exactly 1.0 — a softmax tie.
20203        let lo = -0.1f32;
20204        let hi = f32::from_bits(lo.to_bits() - 1);
20205        assert!(hi > lo && hi - lo < 2f32.powi(-25));
20206        assert_eq!((lo - hi).exp(), 1.0, "the tie this test is about");
20207        // The gated MoE (softmax) path: the tie hands the token to index 0.
20208        let (idx, _, _) = moe_route(&[lo, hi], &moe(false, true), None);
20209        assert_eq!(idx, vec![0], "softmax tie → lower index (the gated contract)");
20210        // The resonance path: the strictly larger raw score wins, weight
20211        // exactly 1.0 with and without norm_topk.
20212        for norm in [true, false] {
20213            let m = moe(true, norm);
20214            let (idx, p, wsum) = moe_route(&[lo, hi, f32::NEG_INFINITY], &m, None);
20215            assert_eq!(idx, vec![1], "norm_topk {norm}");
20216            assert_eq!(p, vec![0.0, 1.0, 0.0]);
20217            assert_eq!(p[1] / wsum, 1.0);
20218            // An exact tie: the first maximum (as `resonance_winner` and
20219            // `embryo_core_route_pick`).
20220            let (idx, _, _) = moe_route(&[hi, hi, lo], &m, None);
20221            assert_eq!(idx, vec![0]);
20222            // `−∞` never wins; the admitted set is honoured.
20223            let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, lo, hi], &m, None);
20224            assert_eq!(idx, vec![2]);
20225            let (idx, p, wsum) = moe_route(&[lo, hi, hi], &m, Some(&[true, false, false]));
20226            assert_eq!(idx, vec![0]);
20227            assert_eq!(p[0] / wsum, 1.0);
20228            // Every admitted expert at −∞: the generic path's uniform
20229            // fallback (lowest index, finite weights).
20230            let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 3], &m, None);
20231            assert_eq!(idx, vec![0]);
20232            assert!(p.iter().all(|v| v.is_finite()) && wsum.is_finite() && wsum > 0.0);
20233        }
20234    }
20235}