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    /// Sliding-window tail trimming, `(slack, align)` for
353    /// `LayerKvCache::trim_window`: every layer with a window keeps only
354    /// the rows its window can still read (Spark-X2.5: 576..639 rows per
355    /// 512-window layer instead of the whole context — 27 of the 4B's 36
356    /// layers). Set at load for Spark-X2.5 only (`CMF_SWA_TRIM=0` turns
357    /// it off); None keeps every row, as every other model always has.
358    pub swa_trim: Option<(usize, usize)>,
359    /// Explicit local/global schedule for architectures that cannot be
360    /// represented by Gemma's every-Nth-global convention.
361    pub sliding_layers: Option<Vec<bool>>,
362    /// Natively bounded anchor record (`arch.anchor_core`): the file's
363    /// operator, installed once at load. `Some` = the model is
364    /// bounded-native — `--o1`/`CMF_O1*` are refused, prefix reuse and
365    /// the penalty window are bounded, and no anchor layer stores
366    /// anything per position.
367    pub anchor_core: Option<cortiq_core::AnchorCoreConfig>,
368    /// The `[W][rd/2]` relative-rotation table every bounded layer shares
369    /// (built from `inv_freq` once the RoPE setup is final).
370    bounded_rope: Option<std::sync::Arc<crate::bounded::BoundedRope>>,
371    /// Bounded prefix-reuse key (length + rolling hash + tail) of a
372    /// bounded-native model; `kv_history` stays empty there so nothing
373    /// in the pipeline grows with the dialogue.
374    pub kv_prefix: KvPrefix,
375    /// Prompt positions the last `generate*` call actually forwarded
376    /// (prefix reuse subtracts what the cache already held).
377    pub last_prefill_tokens: usize,
378    /// RoPE table of the sliding (local) layers, when they use their
379    /// own base frequency (Gemma-3: 10k local vs 1M global).
380    pub inv_freq_local: Option<std::sync::Arc<Vec<f32>>>,
381    pub rotary_dim_local: Option<usize>,
382    pub rope_scale: f32,
383    pub rope_scale_local: f32,
384    /// Gemma-4: global layers run their own geometry — (head_dim,
385    /// num_kv_heads); sliding layers keep the base fields.
386    pub global_attn: Option<(usize, usize)>,
387    /// Gemma-4: the global layers' proportional RoPE table (len
388    /// global_head_dim/2, zero-padded tail = identity rotation).
389    pub inv_freq_global: Option<std::sync::Arc<Vec<f32>>>,
390    /// Scale-less RMS normalization of V heads before caching (Gemma-4).
391    pub attn_v_norm: bool,
392    /// HunYuan dense: per-head q/k norm runs after RoPE (see the arch flag).
393    pub qk_norm_after_rope: bool,
394    /// The per-head `self_attn.g_proj` output gate uses sigmoid
395    /// (Spark-X2.5) instead of softplus (Laguna).
396    pub proj_gate_sigmoid: bool,
397    /// Final-logit soft-capping C: logits = C·tanh(logits/C) (Gemma-4).
398    pub final_softcap: Option<f32>,
399    /// Cortiq Embryo hierarchical head: cluster matrix [C, hidden]. The
400    /// flat logits h·Eᵀ are turned into the two-level log-probabilities
401    /// log softmax_c(h·Cᵀ)[c(v)] + log softmax_{s∈c(v)}(h·E_c(v)ᵀ)[v].
402    pub head_clusters: Option<std::sync::Arc<Vec<f32>>>,
403    /// Gemma-2 attention-logit soft-capping (0.0 = off).
404    pub attn_softcap: f32,
405    /// Compute per-token confidence (a full-vocab softmax each
406    /// token). On by default; `bench --core` turns it off to match
407    /// llama-bench's core timing.
408    confidence_on: bool,
409    /// Test-only one-shot forward failure, scoped to this pipeline so
410    /// parallel scoring tests cannot consume one another's injection.
411    #[cfg(test)]
412    nll_test_fail_at: Option<usize>,
413    /// Test-only route override; avoids mutating the process-wide
414    /// `CMF_PREFILL` environment variable while forcing the serial path.
415    #[cfg(test)]
416    nll_test_force_serial: bool,
417}
418
419#[cfg(target_os = "macos")]
420impl Drop for Pipeline {
421    fn drop(&mut self) {
422        // the async replay writes into `kv_cache` Vecs about to be freed
423        let _ = crate::gpu_metal::wait_replay();
424        crate::gpu::kv_mirror_drop(self.graph_kv_id);
425    }
426}
427
428#[cfg(not(target_os = "macos"))]
429impl Drop for Pipeline {
430    fn drop(&mut self) {
431        // The wgpu resident Embryo graph owns recurrent/KV buffers keyed by
432        // this pipeline's sequence id.  Release that sequence image when a
433        // pooled pipeline is dropped; model weights stay cached for reuse.
434        crate::gpu::graph_kv_reset(self.graph_kv_id);
435    }
436}
437
438/// Model weights. Matrices are `QTensor` (owned f32 for small models
439/// and tests — bit-identical to the historical paths — or quantized
440/// bytes zero-copy from the CMF mmap for big models). 1-D norms are
441/// always small and stay f32.
442pub struct PipelineWeights {
443    /// Embedding table: [vocab_size, hidden_size]
444    pub embed_tokens: QTensor,
445    /// Per-layer weights
446    pub layers: Vec<LayerWeights>,
447    /// LM head: [vocab_size, hidden_size]
448    pub lm_head: QTensor,
449    /// Final norm: [hidden_size]
450    pub final_norm: Vec<f32>,
451}
452
453/// One transformer layer: shared norms + MLP, attention by kind.
454pub struct LayerWeights {
455    pub input_norm: Vec<f32>,
456    /// The pre-FFN norm (`post_attention_layernorm` classically;
457    /// `pre_feedforward_layernorm` on Gemma-2/3 sandwich layers).
458    pub post_norm: Vec<f32>,
459    /// Gemma-2/3 sandwich: norm applied to the ATTENTION OUTPUT before
460    /// its residual add (`post_attention_layernorm` there).
461    pub attn_out_norm: Option<Vec<f32>>,
462    /// Gemma-4: the whole layer output is multiplied by this scalar.
463    pub layer_scale: Option<f32>,
464    /// Gemma-2/3 sandwich: norm applied to the FFN OUTPUT before its
465    /// residual add (`post_feedforward_layernorm`).
466    pub ffn_out_norm: Option<Vec<f32>>,
467    pub ffn: FfnKind,
468    pub attn: AttnKind,
469}
470
471/// FFN gate activation: SiLU (SwiGLU family) or tanh-GELU (Gemma's
472/// GeGLU). A property of the model, carried on every FFN triple.
473#[derive(Clone, Copy, PartialEq, Debug, Default)]
474pub enum Act {
475    #[default]
476    Silu,
477    GeluTanh,
478    /// Exact erf GELU (HF `hidden_act = "gelu"`: Spark-X2.5).
479    Gelu,
480    /// Kimi-K3 SituAndMul: BOTH halves transform —
481    /// a = β·tanh(g/β)·σ(g), up' = linβ·tanh(u/linβ) (linβ>0), out = a·up'.
482    Situ {
483        beta: f32,
484        linear_beta: f32,
485    },
486}
487
488impl Act {
489    pub fn from_arch(name: &str) -> Self {
490        if name == "gelu_tanh" {
491            Self::GeluTanh
492        } else if name == "gelu" {
493            Self::Gelu
494        } else {
495            Self::Silu
496        }
497    }
498
499    /// Arch-driven constructor (activation name + situ betas).
500    pub fn from_arch_full(arch: &cortiq_core::ModelArch) -> Self {
501        match arch.hidden_act.as_str() {
502            "situ" => Self::Situ {
503                beta: arch.activation_situ_beta.unwrap_or(1.0) as f32,
504                linear_beta: arch.activation_situ_linear_beta.unwrap_or(0.0) as f32,
505            },
506            other => Self::from_arch(other),
507        }
508    }
509
510    #[inline]
511    pub fn apply(self, x: f32) -> f32 {
512        match self {
513            Self::Silu => inference::silu(x),
514            Self::GeluTanh => inference::gelu_tanh(x),
515            Self::Gelu => inference::gelu_erf(x),
516            Self::Situ { beta, .. } => beta * (x / beta).tanh() * (1.0 / (1.0 + (-x).exp())),
517        }
518    }
519
520    /// Gated combine — the FFN contract. Situ transforms the UP half
521    /// too, so callers must use this instead of apply(g)·u.
522    #[inline]
523    pub fn combine(self, g: f32, u: f32) -> f32 {
524        match self {
525            Self::Situ { linear_beta, .. } if linear_beta > 0.0 => {
526                self.apply(g) * (linear_beta * (u / linear_beta).tanh())
527            }
528            _ => self.apply(g) * u,
529        }
530    }
531
532    /// The wgpu graphs' arm for this activation; None = no device kernel
533    /// computes it, and the graph builder must refuse the layer (the dense
534    /// graph FFN once computed SiLU for whatever the model asked for).
535    pub fn graph_act(self) -> Option<crate::gpu::GraphAct> {
536        match self {
537            Self::Silu => Some(crate::gpu::GraphAct::Silu),
538            Self::Gelu => Some(crate::gpu::GraphAct::GeluErf),
539            Self::GeluTanh | Self::Situ { .. } => None,
540        }
541    }
542}
543
544/// Dense gated triple — the FFN of a dense layer or of one expert.
545pub struct DenseFfn {
546    pub gate_proj: QTensor,
547    pub up_proj: QTensor,
548    pub down_proj: QTensor,
549    /// Gate activation (SiLU default; Gemma: tanh-GELU).
550    pub act: Act,
551    /// `down_proj` stored transposed (`[inter, hidden]`), when the file
552    /// carries it. Only the per-token sparse path reads it: a neuron's
553    /// down weights are a contiguous ROW there, so the token's chosen
554    /// neurons are the only bytes touched. `None` = the ordinary layout,
555    /// and the sparse path stays off.
556    pub down_t: Option<QTensor>,
557    /// Task tubes (spec: defragged task-conditional width). The three
558    /// matrices above are the CORE — the neurons every task computes;
559    /// each tube is an independently quantized slice of the SAME layer
560    /// holding the neurons only some tasks need. A tube is a normal
561    /// tensor triple, so every kernel runs it unchanged, and the bytes
562    /// of an inactive tube are never read. Empty = ordinary dense FFN.
563    pub segs: Vec<FfnSeg>,
564}
565
566/// One task tube: a contiguous slice of a layer's FFN neurons, stored
567/// as its own `[w, hidden]` / `[hidden, w]` triple. `start` is the
568/// neuron's index in the layer's FULL space (core first, then tubes in
569/// order) — the bit a task mask sets to switch this tube on.
570pub struct FfnSeg {
571    pub gate: QTensor,
572    pub up: QTensor,
573    pub down: QTensor,
574    pub start: usize,
575    pub width: usize,
576}
577
578/// FFN operator of a layer, decided by tensor presence at load time
579/// (router `mlp.gate.weight` in the directory = MoE layer).
580pub enum FfnKind {
581    Dense(DenseFfn),
582    /// Mixture-of-Experts (Qwen2-MoE / Qwen3-MoE): softmax over ALL
583    /// expert logits → top-k, optional renorm; experts stay quantized
584    /// in mmap — only the selected ones are touched per token.
585    Moe(MoeFfn),
586    /// Gemma-4 MoE: a dense MLP branch AND a routed-expert branch in
587    /// the SAME layer, each with its own norm sandwich. The dense
588    /// branch reads the pre-FFN-normed input; the expert branch (and
589    /// the router) read the RAW residual through `pre_norm_2`:
590    ///   d = post_norm_1(dense(x̂));  m = post_norm_2(Σwₑ·FFNₑ(pre_norm_2(h)))
591    ///   ffn_out = d + m   (the caller's ffn_out_norm + residual follow)
592    DenseMoe(Box<DenseMoeFfn>),
593}
594
595/// Gemma-4 dual-branch FFN (see `FfnKind::DenseMoe`).
596pub struct DenseMoeFfn {
597    pub dense: DenseFfn,
598    pub moe: MoeFfn,
599    /// post_feedforward_layernorm_1 — dense-branch output norm.
600    pub post_norm_1: Vec<f32>,
601    /// pre_feedforward_layernorm_2 — expert-branch input norm (applied
602    /// to the RAW residual, not the pre-FFN-normed activation).
603    pub pre_norm_2: Vec<f32>,
604    /// post_feedforward_layernorm_2 — expert-branch output norm.
605    pub post_norm_2: Vec<f32>,
606}
607
608pub struct MoeFfn {
609    /// Router `mlp.gate.weight` [num_experts, hidden].
610    pub router: QTensor,
611    pub experts: Vec<DenseFfn>,
612    pub top_k: usize,
613    pub norm_topk_prob: bool,
614    /// Router scores per-expert with a sigmoid (LFM2-MoE / DeepSeek-V3
615    /// `noaux_tc`) instead of a softmax over all experts (Qwen).
616    pub router_sigmoid: bool,
617    /// Per-expert selection bias `mlp.expert_bias` [num_experts]
618    /// (LFM2-MoE): added to the sigmoid scores for the top-k CHOICE only;
619    /// the gathered weights use the unbiased scores. None = no bias.
620    pub expert_bias: Option<Vec<f32>>,
621    /// Top-k weights are multiplied by this after the optional renorm
622    /// (LFM2-MoE `routed_scaling_factor`; 1.0 = off).
623    pub routed_scaling: f32,
624    /// Adaptive routing (CMF_MOE_TAU, opt-in): keep the smallest
625    /// prefix of the top-k whose renormalized mass reaches τ —
626    /// confident tokens touch 1–2 experts, flat ones keep all k.
627    /// MoE decode is memory-bound, so skipped experts are skipped
628    /// weight traffic. None = classic fixed top-k (bit-identical).
629    pub route_tau: Option<f32>,
630    /// Always-on shared expert. Qwen2-MoE carries an additional sigmoid
631    /// gate; Laguna adds the shared expert unconditionally (`None`).
632    pub shared: Option<(DenseFfn, Option<QTensor>)>,
633    /// Expert-selection counters (truncated Fisher B-field of claim 12:
634    /// routing frequency during calibration). Filled by every forward,
635    /// read by the CLI via CMF_MOE_STATS. RefCell: decode is single-threaded.
636    pub stats: std::cell::RefCell<Vec<u64>>,
637    /// Per-CHANNEL sum of squares of this FFN's input, accumulated over a
638    /// calibration run (`CMF_RMS_TRACE`). These are the RMS activation
639    /// traces AWNP needs: raw weight magnitude says every channel matters
640    /// equally, and the question AWNP asks is whether the ACTIVATIONS
641    /// disagree. Off unless the env var is set — an f64 add per channel
642    /// per token is cheap, but not free.
643    pub act_sq: std::cell::RefCell<Vec<f64>>,
644    /// Raw FFN-input rows captured for the layers named by `CMF_ACT_DUMP`
645    /// (`"9,19"`). AWNP is nullspace PROJECTION: after dropping channels the
646    /// survivors are refitted to absorb what was removed, and how much they
647    /// can absorb depends on the activation COVARIANCE, not on per-channel
648    /// RMS. Per-channel numbers can only bound the cost from above.
649    pub act_rows: std::cell::RefCell<Vec<f32>>,
650    /// Task mask over routed experts (DTG-MA over MoE, claim-12 B-field
651    /// applied): `false` experts are excluded from selection, the
652    /// softmax renormalizes over the allowed set. Built by the loader
653    /// from CMF_MOE_MASK=<stats.json> + CMF_MOE_MASK_COVER. None = all.
654    pub mask: Option<Vec<bool>>,
655    /// Gemma-4: per-expert weight scale applied AFTER the top-k renorm
656    /// (`router.per_expert_scale`). None = 1.0 everywhere.
657    pub per_expert_scale: Option<Vec<f32>>,
658    /// Gemma-4: the router reads a SCALE-LESS rms-norm of its input
659    /// (the constant gain router.scale·√hidden is folded into the
660    /// router weights at convert time).
661    pub router_input_norm: bool,
662    /// Cortiq Embryo: resonance routing (P1) — the "logits" are
663    /// bias_e − ‖(x−μ_e) − U_eᵀU_e(x−μ_e)‖², argmax = the expert whose
664    /// descriptor reconstructs the input best. `router` is a placeholder.
665    pub resonance: Option<Resonance>,
666    /// Growth records (`kind = "expert_append"`, spec §9.5.1) mounted
667    /// behind the trunk experts, one entry per grown expert IN THE ORDER
668    /// they sit in `experts` (the tail `experts[experts.len() - grown.len()..]`).
669    /// The trunk keeps `experts.len() - grown.len()` experts. Empty on a
670    /// gated MoE, on a file without records and under `CMF_GROWTH=off`.
671    pub grown: Vec<GrownExpert>,
672}
673
674/// One grown expert of an `expert_append` record as the loader mounted it.
675#[derive(Debug, Clone, PartialEq, Eq)]
676pub struct GrownExpert {
677    /// The record's skill id.
678    pub record: String,
679    /// Position of the record in `header.skills`.
680    pub record_index: usize,
681    pub layer: usize,
682    /// The index the tensor name declares (`experts.{e}` — the chain rule
683    /// of the format); the executed position may be smaller when an
684    /// earlier record of the layer is not mounted.
685    pub expert: usize,
686}
687
688/// Per-expert resonance descriptors of one MoE layer (`mlp.desc.*`).
689pub struct Resonance {
690    /// [E, hidden]
691    pub mu: Vec<f32>,
692    /// [E, k, hidden] orthonormal directions (k may be 0)
693    pub u: Vec<f32>,
694    pub k: usize,
695    /// [E] selection bias (loss-free balancing, trained online)
696    pub bias: Vec<f32>,
697    /// [E] reconstruction-error shell: an expert whose error
698    /// `‖(x−μ)⊥U‖² = d² − proj` exceeds its shell scores `−∞` (it never
699    /// wins). `+inf` = no shell — every trunk expert; a grown expert
700    /// (`expert_append`) carries the finite `desc.shell` its record stores.
701    /// Empty = no shell anywhere (legacy constructors).
702    pub shell: Vec<f32>,
703}
704
705/// `CMF_GROWTH_SHELL` state: 0 = not yet read from the environment, 1 = on,
706/// 2 = off. Process-wide, like the environment it mirrors.
707static GROWTH_SHELL: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
708
709/// Is the growth shell applied (`Resonance::scores` −∞ rule, the resident
710/// graph's packed shell)? `CMF_GROWTH_SHELL=off` disables it for
711/// measurement; [`set_growth_shell`] overrides the environment in-process
712/// (`growth-eval --shell`). Default: on.
713pub fn growth_shell_enabled() -> bool {
714    use std::sync::atomic::Ordering;
715    match GROWTH_SHELL.load(Ordering::Relaxed) {
716        1 => true,
717        2 => false,
718        _ => {
719            let off = std::env::var("CMF_GROWTH_SHELL")
720                .map(|v| v.eq_ignore_ascii_case("off") || v == "0")
721                .unwrap_or(false);
722            GROWTH_SHELL.store(if off { 2 } else { 1 }, Ordering::Relaxed);
723            !off
724        }
725    }
726}
727
728/// Switch the growth shell on/off for this process (`None` = re-read
729/// `CMF_GROWTH_SHELL` on the next query). A pipeline packed into the
730/// resident graph BEFORE the switch keeps the shell it was packed with —
731/// build a new pipeline after switching.
732pub fn set_growth_shell(on: Option<bool>) {
733    GROWTH_SHELL.store(
734        match on {
735            Some(true) => 1,
736            Some(false) => 2,
737            None => 0,
738        },
739        std::sync::atomic::Ordering::Relaxed,
740    );
741}
742
743impl Resonance {
744    /// Does any expert carry a finite shell (a mounted growth record)?
745    pub fn has_shell(&self) -> bool {
746        self.shell.iter().any(|s| s.is_finite())
747    }
748
749    /// The shell the runtime applies right now: the stored one, or all
750    /// `+inf` when the shell is switched off (`CMF_GROWTH_SHELL=off`).
751    pub fn effective_shell(&self, ne: usize) -> Vec<f32> {
752        let mut out = vec![f32::INFINITY; ne];
753        if growth_shell_enabled() {
754            for (o, s) in out.iter_mut().zip(&self.shell) {
755                *o = *s;
756            }
757        }
758        out
759    }
760
761    /// Routing scores for one input row (higher = better). A grown
762    /// expert whose reconstruction error lies outside its shell gets
763    /// `−∞` (unless the shell is switched off); trunk rows are the exact
764    /// bit pattern they were before growth.
765    pub fn scores(&self, x: &[f32], out: &mut [f32]) {
766        let h = x.len();
767        let ne = out.len();
768        let shell_on = growth_shell_enabled() && !self.shell.is_empty();
769        for e in 0..ne {
770            let mu = &self.mu[e * h..(e + 1) * h];
771            let mut d2 = 0.0f32;
772            for j in 0..h {
773                let d = x[j] - mu[j];
774                d2 += d * d;
775            }
776            let mut proj = 0.0f32;
777            for i in 0..self.k {
778                let u = &self.u[(e * self.k + i) * h..(e * self.k + i + 1) * h];
779                let mut p = 0.0f32;
780                for j in 0..h {
781                    p += (x[j] - mu[j]) * u[j];
782                }
783                proj += p * p;
784            }
785            let err = d2 - proj;
786            out[e] = self.bias.get(e).copied().unwrap_or(0.0) - err;
787            if shell_on && err > self.shell.get(e).copied().unwrap_or(f32::INFINITY) {
788                out[e] = f32::NEG_INFINITY;
789            }
790        }
791    }
792}
793
794/// Attention operator of a layer. Extension point: new operators are
795/// new variants here + a forward in their own module.
796pub enum AttnKind {
797    /// GQA softmax attention (+ optional Qwen3.5 qk-norm / output gate).
798    Full {
799        wq: QTensor,
800        wk: QTensor,
801        wv: QTensor,
802        wo: QTensor,
803        q_norm: Option<Vec<f32>>,
804        k_norm: Option<Vec<f32>>,
805        output_gate: bool,
806        /// Laguna: a separate softplus projection applied to the attention
807        /// output before O. The bool means one scalar per head (broadcast
808        /// across head_dim); false means one scalar per element.
809        softplus_gate: Option<(QTensor, bool)>,
810        /// Qwen2-family projection biases (q, k, v).
811        bias: Option<(Vec<f32>, Vec<f32>, Vec<f32>)>,
812    },
813    /// Canonical linear core (VMF phase attention).
814    Linear(VmfPhaseWeights),
815    /// Faithful vendor linear operator (Qwen3.5 GatedDeltaNet).
816    LinearGdn(GdnWeights),
817    /// LFM2 gated short-convolution mixer (no KV cache; conv ring state
818    /// lives in the layer's `linear_state`).
819    ShortConv(ShortConvWeights),
820    /// DeepSeek-V2 Multi-head Latent Attention. v1 executes it as
821    /// expand-to-MHA: the latent is projected per token, K/V expand to
822    /// every head and live in the ordinary cache (K head layout
823    /// [rope | nope] so the standard partial rotary covers the shared
824    /// rope key; V rows are zero-padded to the K head_dim and the pad
825    /// is sliced off before O). Latent-resident cache is a later
826    /// optimization, not a semantic change.
827    Mla(Box<MlaWeights>),
828    /// Kimi Delta Attention (Kimi Linear / Kimi-K3): per-channel decayed
829    /// delta rule, separate q/k/v short convs, sigmoid-gated output norm.
830    /// State lives in the layer's `linear_state` (no KV cache).
831    Kda(Box<crate::linear_core::KdaWeights>),
832    /// Natively bounded softmax anchor `swa_sink_v1` (Embryo-O1): ring of
833    /// the last W raw keys with relative RoPE + trained NoPE sinks, one
834    /// softmax. State is the fixed-size ring in `LayerKvCache::bounded`;
835    /// nothing is stored per position (see `crate::bounded`).
836    Bounded(Box<crate::bounded::BoundedWeights>),
837}
838
839/// DeepSeek-V2 MLA projections (see `AttnKind::Mla`).
840pub struct MlaWeights {
841    /// `[nh·(rope+nope), hidden]` (or `[…, q_lora]` when compressed) —
842    /// the converter permutes each head rope-first so rotary_dim =
843    /// qk_rope works unchanged.
844    pub q_proj: QTensor,
845    /// Compressed q (K3/V3 class): x → q_a `[q_lora, hidden]` →
846    /// rms(q_a_norm) → q_proj (= q_b). None = direct q (V2-Lite).
847    pub q_a: Option<QTensor>,
848    pub q_a_norm: Option<Vec<f32>>,
849    /// `kv_a_proj_with_mqa` `[lora + rope, hidden]` (latent first).
850    pub kv_a: QTensor,
851    /// RMS-norm weights over the latent (`kv_a_layernorm`, [lora]).
852    pub kv_a_norm: Vec<f32>,
853    /// `[nh·(nope+v), lora]` — per head [k_nope | v].
854    pub kv_b: QTensor,
855    /// `[hidden, nh·v]`.
856    pub o_proj: QTensor,
857    pub nh: usize,
858    pub qk_rope: usize,
859    pub qk_nope: usize,
860    pub v_dim: usize,
861    pub lora: usize,
862    /// Softmax scale (1/√(rope+nope), YaRN-mscale-corrected at load).
863    pub scale: f32,
864    /// Kimi Linear NoPE: skip the rotary entirely (layout unchanged).
865    pub nope: bool,
866}
867
868/// Multi-token-prediction head (DeepSeek/Qwen style, spec §2.1):
869/// `x = eh_proj·[enorm(embed(next)); hnorm(hidden)]` → one transformer
870/// block over its own KV → shared lm_head. Drafts the token after next;
871/// the main model verifies, so output is exact — MTP only buys speed.
872pub struct MtpModule {
873    pub enorm: Vec<f32>,
874    pub hnorm: Vec<f32>,
875    /// [hidden, 2·hidden]
876    pub eh_proj: QTensor,
877    pub layer: LayerWeights,
878    pub final_norm: Vec<f32>,
879    pub kv: crate::kv_cache::LayerKvCache,
880}
881
882/// A Metal verify graph after its sync: what the commit needs — the
883/// graph (per-layer replay scratch), the GDN layers in encode order (their
884/// CPU states receive the replay), and the attention layers with the CPU
885/// row count they were encoded against (the accepted rows are pulled from
886/// the mirror from there).
887/// One item of the Metal rows-graph plan.
888#[cfg(target_os = "macos")]
889enum MetalRowsItem<'a> {
890    Gdn {
891        run: Vec<crate::gpu_metal::GdnGpuLayer<'a>>,
892        first: usize,
893    },
894    Attn {
895        l: crate::gpu_metal::AttnGpuLayer<'a>,
896        li: usize,
897        q_norm: Option<&'a [f32]>,
898        k_norm: Option<&'a [f32]>,
899        output_gate: bool,
900    },
901}
902
903#[cfg(target_os = "macos")]
904struct MetalVerifyPending {
905    graph: crate::gpu_metal::VerifyGraph,
906    gdn_layers: Vec<usize>,
907    attn_layers: Vec<(usize, usize)>,
908}
909
910/// A round's batched MTP warm-up, submitted but not yet waited
911/// (`mtp_warm_batch_submit` → `mtp_warm_batch_finish`): the trunk commit's
912/// GDN replay is queued between the two.
913#[cfg(target_os = "macos")]
914struct MetalWarmPending {
915    graph: crate::gpu_metal::VerifyGraph,
916    cpu_stored: usize,
917    b: usize,
918}
919
920#[cfg(target_os = "macos")]
921enum MetalRowsRun {
922    /// Capability/preflight refusal before a command buffer was committed.
923    Declined,
924    /// A graph was admitted and then failed; callers must clear the sequence
925    /// rather than replaying it through CPU/serial state.
926    Failed,
927    Completed(MetalVerifyPending),
928}
929
930#[cfg(target_os = "macos")]
931enum MetalPrefillOutcome {
932    Declined,
933    Failed,
934    Completed(Vec<f32>),
935}
936
937#[cfg(target_os = "macos")]
938enum MetalBatchNllOutcome {
939    Declined,
940    Failed(String),
941    Completed(f64, usize),
942}
943
944/// The speculation trial's phases (see the decode loop): four timed
945/// speculative rounds, eight timed plain tokens, then the faster arm
946/// until a re-check.
947#[derive(Clone, Copy)]
948enum SpecTrial {
949    Spec {
950        t0: std::time::Instant,
951        gen0: usize,
952        rounds: usize,
953    },
954    Plain {
955        t0: std::time::Instant,
956        gen0: usize,
957    },
958    Decided {
959        spec: bool,
960        recheck_at: usize,
961    },
962}
963
964/// `CMF_GRAPH_SPEC_TIME`: 0 = off, 1 = one line per speculative round
965/// plus the host stamps of any OUTLIER round (wall > 1.4× the running
966/// median), 2 = the host stamps of every round.
967pub(crate) fn spec_time_level() -> u8 {
968    static L: std::sync::OnceLock<u8> = std::sync::OnceLock::new();
969    *L.get_or_init(|| match std::env::var("CMF_GRAPH_SPEC_TIME") {
970        Ok(v) => v.trim().parse::<u8>().map(|n| n.max(1)).unwrap_or(1),
971        Err(_) => 0,
972    })
973}
974
975/// The round's host stamps: `spec_stamp(name)` records the time since
976/// the previous stamp (the section that just ended) — from anywhere on
977/// the round's call chain (the Metal verify, the draft step, the commit),
978/// no plumbing. Off (a single atomic load) unless `CMF_GRAPH_SPEC_TIME`
979/// is set; one decode thread at a time is assumed (diagnostics).
980struct SpecStampLog {
981    t_last: std::time::Instant,
982    items: Vec<(&'static str, f32)>,
983}
984
985static SPEC_STAMPS: std::sync::Mutex<Option<SpecStampLog>> = std::sync::Mutex::new(None);
986
987pub(crate) fn spec_stamp(name: &'static str) {
988    if spec_time_level() == 0 {
989        return;
990    }
991    if let Ok(mut g) = SPEC_STAMPS.lock() {
992        if let Some(log) = g.as_mut() {
993            let now = std::time::Instant::now();
994            log.items
995                .push((name, (now - log.t_last).as_secs_f32() * 1e3));
996            log.t_last = now;
997        }
998    }
999}
1000
1001fn spec_stamps_begin() {
1002    if spec_time_level() == 0 {
1003        return;
1004    }
1005    if let Ok(mut g) = SPEC_STAMPS.lock() {
1006        *g = Some(SpecStampLog {
1007            t_last: std::time::Instant::now(),
1008            items: Vec::with_capacity(64),
1009        });
1010    }
1011}
1012
1013fn spec_stamps_take() -> Vec<(&'static str, f32)> {
1014    SPEC_STAMPS
1015        .lock()
1016        .ok()
1017        .and_then(|mut g| g.take())
1018        .map(|l| l.items)
1019        .unwrap_or_default()
1020}
1021
1022/// One line: every stamp name in first-seen order with its total over the
1023/// round and, when it fired more than once (the draft steps), the count.
1024fn spec_stamps_format(items: &[(&'static str, f32)]) -> String {
1025    let mut agg: Vec<(&'static str, f32, u32)> = Vec::with_capacity(items.len());
1026    for &(n, ms) in items {
1027        match agg.iter_mut().find(|e| e.0 == n) {
1028            Some(e) => {
1029                e.1 += ms;
1030                e.2 += 1;
1031            }
1032            None => agg.push((n, ms, 1)),
1033        }
1034    }
1035    let mut s = String::with_capacity(agg.len() * 16);
1036    for (n, ms, k) in agg {
1037        if k > 1 {
1038            s.push_str(&format!("{n} {ms:.1}/{k} "));
1039        } else {
1040            s.push_str(&format!("{n} {ms:.1} "));
1041        }
1042    }
1043    s
1044}
1045
1046/// The speculation monitor: exponential averages of a round's wall time
1047/// and of the tokens it produced, and the plain token's wall time — the
1048/// three numbers the keep/stop rule needs. A round pays when
1049/// `tokens_per_round · plain_ms > round_ms · 1.03`. The one-shot trial
1050/// (four rounds against eight tokens) mis-called prose: the first rounds
1051/// after a prompt are formulaic and accept well, the body does not (an
1052/// essay measured 39 against a plain 44.8 with the trial saying
1053/// "speculate"), so the rule now runs on EVERY round and stops after four
1054/// consecutive losing rounds; a stopped speculation is retried 128 tokens
1055/// later.
1056///
1057/// Native Metal (`metal: true`) does not pay the eight plain tokens up
1058/// front: on the 27B a plain token is ~150 ms, so the trial alone cost
1059/// ~1.2 s of every answer. There the plain phase is (a) skipped while the
1060/// rounds land at least `SPEC_PROXY_TOKENS` tokens each — a k=7 round on
1061/// Metal costs ~1.9 plain tokens (286 against 148 ms measured on the M4),
1062/// so 3.5 tokens/round cannot lose on any Metal round/plain ratio seen —
1063/// and (b) otherwise bounded to the fewest tokens that time it: two, or
1064/// as many as fit in `SPEC_PLAIN_MIN_MS` (a 150-ms token measures itself;
1065/// a 10-ms one needs the eight). The keep/stop rule itself is unchanged:
1066/// the moment a plain rate exists, it decides.
1067#[derive(Default, Clone, Copy)]
1068struct SpecMon {
1069    round_ms: f64,
1070    tokens: f64,
1071    plain_ms: f64,
1072    n: u32,
1073    fails: u32,
1074    metal: bool,
1075}
1076
1077/// Tokens per round at or above which a Metal round pays without a plain
1078/// measurement (see `SpecMon`).
1079const SPEC_PROXY_TOKENS: f64 = 3.5;
1080/// The Metal plain phase: at least two tokens, and more until this much
1081/// wall time has been timed (up to the eight the other backends time).
1082const SPEC_PLAIN_MIN_MS: f64 = 200.0;
1083
1084impl SpecMon {
1085    fn round(&mut self, dt_ms: f64, produced: usize) {
1086        self.n += 1;
1087        if self.n == 1 {
1088            return; // round 1 pays the batch scratch and the draft mirror
1089        }
1090        let a = if self.n == 2 { 1.0 } else { 0.3 };
1091        self.round_ms += a * (dt_ms - self.round_ms);
1092        self.tokens += a * (produced as f64 - self.tokens);
1093    }
1094    fn pays(&self) -> bool {
1095        if self.plain_ms > 0.0 {
1096            self.tokens * self.plain_ms > self.round_ms * 1.03
1097        } else {
1098            self.metal && self.tokens >= SPEC_PROXY_TOKENS
1099        }
1100    }
1101    /// Has the plain phase timed enough tokens to decide?
1102    fn plain_done(&self, t0: std::time::Instant, gen0: usize, generated: usize) -> bool {
1103        let n = generated.saturating_sub(gen0);
1104        if n >= 8 {
1105            return true;
1106        }
1107        self.metal && n >= 2 && t0.elapsed().as_secs_f64() * 1e3 >= SPEC_PLAIN_MIN_MS
1108    }
1109}
1110
1111/// Ids of the consumed prefix the bounded reuse key remembers literally
1112/// (the rest is covered by the rolling hash).
1113pub const KV_PREFIX_TAIL: usize = 128;
1114
1115/// Bounded prefix-reuse key: how many ids the cache holds, a rolling
1116/// hash of ALL of them and the last [`KV_PREFIX_TAIL`] ids literally.
1117/// Answers "does this prompt strictly extend what is cached" exactly
1118/// (hash over the whole consumed prefix + literal tail) without keeping
1119/// the dialogue — the record is a fixed size whatever the session length.
1120#[derive(Debug, Clone, Default)]
1121pub struct KvPrefix {
1122    len: usize,
1123    hash: u64,
1124    tail: Vec<u32>,
1125    /// Owner of the state this key describes: the resident device graph
1126    /// (true) or the host. A turn continues the prefix only on its owner
1127    /// (R4: a host continuation of a device sequence reads an empty host
1128    /// state; a device continuation of a host sequence has no image).
1129    device: bool,
1130}
1131
1132impl KvPrefix {
1133    #[inline]
1134    fn fold(mut h: u64, ids: &[u32]) -> u64 {
1135        for &id in ids {
1136            h ^= id as u64;
1137            h = h.wrapping_mul(0x100000001b3);
1138            h ^= h >> 29;
1139        }
1140        h
1141    }
1142
1143    pub fn clear(&mut self) {
1144        self.len = 0;
1145        self.hash = 0xcbf29ce484222325;
1146        self.tail.clear();
1147        self.device = false;
1148    }
1149
1150    /// Was the prefix built on the resident device graph?
1151    pub fn on_device(&self) -> bool {
1152        self.device
1153    }
1154
1155    /// Tag the owner of the recorded prefix.
1156    pub fn set_on_device(&mut self, device: bool) {
1157        self.device = device;
1158    }
1159
1160    /// Ids the cache holds (the forwarded prefix).
1161    pub fn len(&self) -> usize {
1162        self.len
1163    }
1164
1165    pub fn is_empty(&self) -> bool {
1166        self.len == 0
1167    }
1168
1169    /// Literal tail currently kept (≤ `KV_PREFIX_TAIL`).
1170    pub fn tail_len(&self) -> usize {
1171        self.tail.len()
1172    }
1173
1174    /// Replace the key with `ids` (a fresh sequence).
1175    pub fn set(&mut self, ids: &[u32]) {
1176        self.clear();
1177        self.extend(ids);
1178    }
1179
1180    /// Append `more` to the consumed prefix (an extension-only turn).
1181    pub fn extend(&mut self, more: &[u32]) {
1182        if self.len == 0 && self.hash == 0 {
1183            self.hash = 0xcbf29ce484222325;
1184        }
1185        self.hash = Self::fold(self.hash, more);
1186        self.len += more.len();
1187        if more.len() >= KV_PREFIX_TAIL {
1188            self.tail.clear();
1189            self.tail.extend_from_slice(&more[more.len() - KV_PREFIX_TAIL..]);
1190        } else {
1191            let drop = (self.tail.len() + more.len()).saturating_sub(KV_PREFIX_TAIL);
1192            self.tail.drain(..drop);
1193            self.tail.extend_from_slice(more);
1194        }
1195    }
1196
1197    /// Cached positions when `ids` strictly extends the consumed prefix,
1198    /// 0 otherwise. The tail is compared literally first (cheap), then
1199    /// the hash over the whole prefix must agree.
1200    pub fn extension(&self, ids: &[u32]) -> usize {
1201        if self.len == 0 || ids.len() <= self.len {
1202            return 0;
1203        }
1204        let t = self.tail.len();
1205        if ids[self.len - t..self.len] != self.tail[..] {
1206            return 0;
1207        }
1208        if Self::fold(0xcbf29ce484222325, &ids[..self.len]) != self.hash {
1209            return 0;
1210        }
1211        self.len
1212    }
1213}
1214
1215/// Result of a generation call.
1216pub struct GenerateResult {
1217    pub text: String,
1218    pub token_ids: Vec<u32>,
1219    pub prompt_tokens: usize,
1220    pub tokens_generated: usize,
1221    pub finish_reason: String,
1222    /// Speculative-decode stats (0/0 when MTP is absent or inactive).
1223    pub mtp_drafted: usize,
1224    pub mtp_accepted: usize,
1225    /// Per-generated-token confidence = softmax probability of the token
1226    /// that was actually emitted (softmax probability on the chosen state). High =
1227    /// the model was sure; low = it was guessing. Same length as the
1228    /// generated slice of `token_ids`.
1229    pub token_confidence: Vec<f32>,
1230    /// Structured per-token telemetry (B4 channel). Empty unless
1231    /// `set_trace(true)`; otherwise same length as the generated slice.
1232    pub traces: Vec<TokenTrace>,
1233}
1234
1235/// One row of the structured telemetry trace (B4): the model's internal
1236/// routing state at the moment a token was emitted. Every field is a
1237/// quantity the runtime already computes — nothing is inferred or
1238/// estimated (anti-principle: only measured bytes).
1239#[derive(Clone, Debug)]
1240pub struct TokenTrace {
1241    /// 0-based index within the generated slice.
1242    pub t: usize,
1243    /// The emitted token id.
1244    pub token_id: u32,
1245    /// Softmax probability on the emitted token — how sure the model was.
1246    pub confidence: f32,
1247    /// Skill in force while this token was generated (None = backbone).
1248    pub active_skill: Option<String>,
1249    /// Recon error E = ‖r−BBᵀr‖²/‖φ‖² at the last routing eval — coherence
1250    /// with the active skill's subspace (low = coherent). None = no router
1251    /// or not yet evaluated.
1252    pub recon: Option<f32>,
1253    /// The router changed the active skill right after this token (a
1254    /// domain boundary crossed under the hysteresis barrier).
1255    pub switched: bool,
1256}
1257
1258/// Calibrated softmax probability of `id` under `logits` (the confidence on
1259/// the emitted token) — the confidence signal, cheap from logits already
1260/// computed for sampling. `temp` is the calibration temperature (B1):
1261/// softmax(logits / temp); 1.0 = raw.
1262#[cfg_attr(not(test), allow(dead_code))]
1263fn top1_prob_t(logits: &[f32], id: u32, temp: f32) -> f32 {
1264    let t = if temp > 1e-3 { temp } else { 1.0 };
1265    let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
1266    let sum: f32 = logits.iter().map(|&v| ((v - max) / t).exp()).sum();
1267    if sum > 0.0 {
1268        (((logits[id as usize] - max) / t).exp()) / sum
1269    } else {
1270        0.0
1271    }
1272}
1273
1274/// prefill-GEMM enabled? (CMF_PREFILL=seq — emergency fallback to the
1275/// sequential path.)
1276fn prefill_batched() -> bool {
1277    std::env::var("CMF_PREFILL")
1278        .map(|v| v != "seq")
1279        .unwrap_or(true)
1280}
1281
1282/// Decide the graph NLL route without conflating graph quality with the
1283/// optional native-Metal fused head. A hidden-state graph remains a valid
1284/// quality route on Vulkan/Wgpu; only native Metal requires graph logits.
1285#[inline]
1286fn nll_graph_policy(
1287    unmasked: bool,
1288    prefer_graph: bool,
1289    native_metal: bool,
1290) -> (bool, bool) {
1291    let graph_quality = unmasked && prefer_graph;
1292    let fused_head_quality = graph_quality && native_metal;
1293    (graph_quality, fused_head_quality)
1294}
1295
1296/// Input to the layer-major batched span walk: token ids (embeds itself,
1297/// full-stack and coordinator prefill) or ready boundary hiddens (the
1298/// network worker's side of a split).
1299#[derive(Clone, Copy)]
1300enum PrefillIn<'a> {
1301    Ids(&'a [u32]),
1302    Hidden(&'a [f32]),
1303}
1304
1305/// The batched prefill walks `weights.layers`. Architectures that load
1306/// their own stack (gemma-3n's AltUp replicas, DeepSeek-V4's hyper-
1307/// connections) leave that empty and must go position by position — asking
1308/// otherwise indexes an empty vector, which is a panic rather than a
1309/// fallback. Every call site goes through here so the next such
1310/// architecture is one line, not four.
1311impl Pipeline {
1312    fn can_prefill_batched(&self) -> bool {
1313        #[cfg(test)]
1314        let force_serial = self.nll_test_force_serial;
1315        #[cfg(not(test))]
1316        let force_serial = false;
1317        prefill_batched() && !force_serial && !self.weights.layers.is_empty()
1318    }
1319
1320    /// The backend's automatic capacity split for a mapped transformer.
1321    /// Kept as a method so prefill and decode use the exact same boundary.
1322    fn automatic_gpu_prefix(&self) -> Option<usize> {
1323        let (model, _, _, _) = self.weights.embed_tokens.graph_weight()?;
1324        crate::gpu::automatic_layer_prefix(&model, self.num_layers, self.physical_layers)
1325    }
1326
1327    /// Positions per batched pass of the layer-stack prefill for THIS
1328    /// model on THIS backend (see [`prefill_chunk_rule`]). Pub: the network
1329    /// split must chunk exactly like the local path to reproduce it.
1330    pub fn prefill_chunk(&self) -> usize {
1331        let env = env_prefill_chunk();
1332        if env.is_some() || ChunkHost::here() != ChunkHost::Other {
1333            return prefill_chunk_rule(env, ChunkHost::here(), false);
1334        }
1335        prefill_chunk_rule(None, ChunkHost::Other, self.chunk_stack_facts().dense_on_discrete())
1336    }
1337
1338    fn chunk_stack_facts(&self) -> ChunkStackFacts {
1339        let plain_dense = !self.weights.layers.is_empty()
1340            && self.g3n.is_none()
1341            && self.dsv4.is_none()
1342            && self.dsv41.is_none()
1343            && self.qwen4_exp.is_none()
1344            && self.weights.layers.iter().all(|lw| {
1345                matches!(lw.attn, AttnKind::Full { .. }) && matches!(lw.ffn, FfnKind::Dense(_))
1346            });
1347        let gpu_on = crate::gpu::enabled();
1348        ChunkStackFacts {
1349            plain_dense,
1350            discrete: gpu_on && crate::gpu::discrete(),
1351            gpu_on,
1352            // Only asked when the rest already qualifies: it opens the
1353            // backend's capacity plan.
1354            capacity_split: std::env::var_os("CMF_GPU_LAYERS").is_some()
1355                || (plain_dense && gpu_on && self.automatic_gpu_prefix().is_some()),
1356            multi_gpu: self.gpu_plan.is_some(),
1357            o1: self.o1_active(),
1358        }
1359    }
1360}
1361
1362/// Prefill chunk (positions per batched pass), model-agnostic form. On
1363/// macOS the AMX GEMM path wants tall panels — M=48 starves the matrix
1364/// units (ggml uses ubatch 512); elsewhere the historical 48 stays.
1365/// CMF_PREFILL_CHUNK overrides. The architectures with their own stacks
1366/// (DeepSeek-V4/V4.1) chunk with this; the layer-stack prefill asks
1367/// [`Pipeline::prefill_chunk`], which also knows the model and the card.
1368/// A different chunk is a different (equally valid) generation: panel
1369/// width reorders float accumulation.
1370pub fn prefill_chunk() -> usize {
1371    prefill_chunk_rule(env_prefill_chunk(), ChunkHost::here(), false)
1372}
1373
1374fn env_prefill_chunk() -> Option<usize> {
1375    std::env::var("CMF_PREFILL_CHUNK")
1376        .ok()
1377        .and_then(|v| v.parse::<usize>().ok())
1378}
1379
1380/// The host classes the chunk width distinguishes.
1381#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1382enum ChunkHost {
1383    Macos,
1384    /// Linux/Android aarch64 (phones, SBCs).
1385    Aarch64,
1386    /// Everything else: x86-64 Linux/Windows, CPU or Vulkan/DX12.
1387    Other,
1388}
1389
1390impl ChunkHost {
1391    fn here() -> Self {
1392        if cfg!(target_os = "macos") {
1393            ChunkHost::Macos
1394        } else if cfg!(target_arch = "aarch64") {
1395            ChunkHost::Aarch64
1396        } else {
1397            ChunkHost::Other
1398        }
1399    }
1400}
1401
1402/// Chunk for a plain dense stack whose every layer lives on a discrete
1403/// card. On x86 the layer-stack prefill is host-driven: each GEMM and the
1404/// chunk attention (which re-uploads the whole KV prefix per layer) is a
1405/// separate submit + readback, so 48 positions a pass left the card idle
1406/// between them. Measured in-process on an RTX 3090 (Vulkan), 2048-token
1407/// prompt — see CHANGELOG 0.7.6 for the table.
1408const DISCRETE_DENSE_PREFILL_CHUNK: usize = 512;
1409
1410/// The chunk-width rule. `dense_on_discrete` is true only for a plain
1411/// dense transformer (full attention, dense FFN, no special stack) that
1412/// is entirely resident on one discrete card — the one case measured
1413/// here. GDN hybrids, MoE, DeepSeek stacks, capacity-split and CPU-only
1414/// runs keep the width they were tuned with.
1415fn prefill_chunk_rule(env: Option<usize>, host: ChunkHost, dense_on_discrete: bool) -> usize {
1416    if let Some(n) = env {
1417        return n.max(1);
1418    }
1419    match host {
1420        ChunkHost::Macos => 512,
1421        // Mobile: big enough to feed the batched attend (gate b ≥ 32)
1422        // and the blocked SDOT GEMM without the memory of 512.
1423        ChunkHost::Aarch64 => 256,
1424        ChunkHost::Other if dense_on_discrete => DISCRETE_DENSE_PREFILL_CHUNK,
1425        ChunkHost::Other => 48,
1426    }
1427}
1428
1429/// What the chunk rule needs to know about a loaded stack.
1430#[derive(Clone, Copy, Debug, Default)]
1431struct ChunkStackFacts {
1432    /// Every layer is `AttnKind::Full` + `FfnKind::Dense`, and no
1433    /// architecture-owned stack (g3n, DeepSeek-V4/V4.1, qwen4-exp) is set.
1434    plain_dense: bool,
1435    /// The active GPU backend is a discrete card.
1436    discrete: bool,
1437    /// The backend is up and not paused.
1438    gpu_on: bool,
1439    /// A capacity-derived device prefix: some layers run on the host.
1440    capacity_split: bool,
1441    /// An in-process multi-GPU plan is set.
1442    multi_gpu: bool,
1443    /// O(1) layers (their Q trace is recorded by the prefill).
1444    o1: bool,
1445}
1446
1447impl ChunkStackFacts {
1448    fn dense_on_discrete(self) -> bool {
1449        self.plain_dense
1450            && self.discrete
1451            && self.gpu_on
1452            && !self.capacity_split
1453            && !self.multi_gpu
1454            && !self.o1
1455    }
1456}
1457
1458/// Number of prompt rows that have a real teacher-forced next-token pair in a
1459/// prefill span.  The final prompt row has no successor token, so it must not
1460/// be handed to the MTP warm-up.  Keeping this arithmetic in one helper makes
1461/// the full-chunk and tail-chunk boundaries explicit for both the graph and
1462/// CPU implementations.
1463#[inline]
1464fn mtp_prefill_pair_count(start: usize, end: usize, input_len: usize) -> usize {
1465    if end <= start || start >= input_len {
1466        return 0;
1467    }
1468    let rows = (end.min(input_len) - start).min(input_len - start);
1469    if end < input_len {
1470        rows
1471    } else {
1472        rows.saturating_sub(1)
1473    }
1474}
1475
1476/// Callback for streaming tokens. Return `false` to cancel.
1477pub type TokenCallback = Box<dyn FnMut(&str) -> bool + Send>;
1478
1479/// One layer's cache ownership at a cross-turn KV reuse boundary.
1480#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1481pub(crate) struct ReuseLayer {
1482    /// Exact-attention layer (rows in `LayerKvCache`); otherwise a
1483    /// recurrent / latent mixer whose state cannot be rewound.
1484    pub full: bool,
1485    /// Positions the host owner cache reaches (`pos_len`: a trimmed
1486    /// sliding tail stores only the newest of them).
1487    pub host_rows: usize,
1488    /// Rows the wgpu token graph's device mirror holds (None: no mirror).
1489    pub device_rows: Option<usize>,
1490    /// A recurrent state lives on the device (advanced past the host copy).
1491    pub device_state: bool,
1492}
1493
1494/// What a reused turn must do before its tail prefill runs on the HOST.
1495#[derive(Debug, Clone, PartialEq, Eq)]
1496pub(crate) enum ReusePlan {
1497    /// Host caches already hold exactly the reused prefix.
1498    Ready,
1499    /// Copy device mirror rows `[from..to)` into the host cache of each
1500    /// listed layer (the rows decode wrote on the device only).
1501    Pull(Vec<(usize, usize, usize)>),
1502    /// The prefix cannot be continued on the host exactly: start fresh.
1503    Fresh,
1504}
1505
1506/// The wgpu whole-token graph decodes into a DEVICE K/V mirror and never
1507/// writes those rows back to the host cache, while the chunked prefill of a
1508/// pure-attention model reads (and appends to) the host cache. A reused turn
1509/// therefore found its host cache ending at the previous PROMPT, not at the
1510/// previous answer: the tail prefill attended without the model's own
1511/// answer and appended its rows at the wrong index (MiniCPM5 on Vulkan
1512/// repeated its tool call instead of reading the tool result). Every layer
1513/// must hold exactly `reuse_from` host rows before the host continues; rows
1514/// that exist only on the device are pulled back, anything else is fresh.
1515pub(crate) fn kv_reuse_plan(reuse_from: usize, layers: &[ReuseLayer]) -> ReusePlan {
1516    let mut pulls = Vec::new();
1517    for (li, l) in layers.iter().enumerate() {
1518        if !l.full {
1519            if l.device_state {
1520                return ReusePlan::Fresh;
1521            }
1522            continue;
1523        }
1524        if l.host_rows == reuse_from {
1525            continue;
1526        }
1527        if l.host_rows < reuse_from && l.device_rows.is_some_and(|d| d >= reuse_from) {
1528            pulls.push((li, l.host_rows, reuse_from));
1529            continue;
1530        }
1531        return ReusePlan::Fresh;
1532    }
1533    if pulls.is_empty() {
1534        ReusePlan::Ready
1535    } else {
1536        ReusePlan::Pull(pulls)
1537    }
1538}
1539
1540impl Pipeline {
1541    /// Clear all per-sequence state, including backend device mirrors.
1542    ///
1543    /// The host KV/history buffers are only half of the request lifecycle on
1544    /// wgpu: GDN/O(1) state and cached graph bind groups are keyed by the
1545    /// pipeline id and otherwise survive a pooled request.  Keep every fresh
1546    /// sequence entry point on this one reset path so a new request cannot
1547    /// inherit the prior request's device state.
1548    fn clear_sequence_state(&mut self) {
1549        // a replay still writing the GDN owners must land before they are
1550        // cleared or reallocated (the device holds raw pointers to them)
1551        #[cfg(target_os = "macos")]
1552        let _ = crate::gpu_metal::wait_replay();
1553        self.kv_cache.clear();
1554        // Both reuse keys (the legacy `kv_history` and the bounded
1555        // `kv_prefix`) describe the state being dropped here.
1556        self.clear_history();
1557        self.graph_logits = None;
1558        if let Some(b) = &mut self.dsv41 {
1559            b.3.clear();
1560        }
1561        crate::gpu::graph_kv_reset(self.graph_kv_id);
1562        // MTP is detached from `self` for the duration of generation, so its
1563        // device mirror is not covered by the trunk reset above.  Reset the
1564        // derived id as well: a failed/aborted warm-up must never leave a
1565        // mirror that a later request can mistake for a current MTP cache.
1566        crate::gpu::graph_kv_reset(self.mtp_kv_id());
1567    }
1568
1569    /// Make the host caches own exactly the reused prefix `[0..reuse_from)`
1570    /// before a reused turn's tail prefill runs on the host (see
1571    /// [`kv_reuse_plan`]). Returns false when the prefix cannot be continued
1572    /// exactly — the caller then starts a fresh sequence. A model whose
1573    /// prefill runs through the token graph keeps its device state as the
1574    /// authority and is left untouched.
1575    fn prepare_kv_reuse(&mut self, reuse_from: usize) -> bool {
1576        if self.graph_prefill_preferred() {
1577            return true;
1578        }
1579        let kv_id = self.graph_kv_id;
1580        let layers: Vec<ReuseLayer> = (0..self.num_layers)
1581            .map(|li| {
1582                let full = matches!(
1583                    self.weights.layers[self.phys_layer(li)].attn,
1584                    AttnKind::Full { .. }
1585                );
1586                ReuseLayer {
1587                    full,
1588                    // Absolute depth: a trimmed sliding tail stores fewer
1589                    // rows than the positions it has seen.
1590                    host_rows: self.kv_cache.layers[li].pos_len(),
1591                    device_rows: crate::gpu::graph_kv_stored(kv_id, li),
1592                    device_state: crate::gpu::graph_state_resident(kv_id, li),
1593                }
1594            })
1595            .collect();
1596        // No wgpu device state at all (CPU, Metal — whose graph appends every
1597        // decoded row to the owner cache itself): the host is the owner and
1598        // the extension check already proved the prefix.
1599        if layers
1600            .iter()
1601            .all(|l| l.device_rows.is_none() && !l.device_state)
1602        {
1603            return true;
1604        }
1605        let plan = kv_reuse_plan(reuse_from, &layers);
1606        let (what, rows, n) = match &plan {
1607            ReusePlan::Ready => ("host ready", 0, 0),
1608            ReusePlan::Fresh => ("fresh", 0, 0),
1609            ReusePlan::Pull(p) => (
1610                "pull",
1611                p.iter().map(|&(_, a, b)| b - a).max().unwrap_or(0),
1612                p.len(),
1613            ),
1614        };
1615        let t0 = std::time::Instant::now();
1616        let ok = self.apply_kv_reuse_plan(reuse_from, plan, &layers);
1617        if std::env::var("CMF_PREFILL_PROF").is_ok() {
1618            eprintln!(
1619                "kv-reuse: {what}{}: {rows} device row(s) × {n} layer(s) to the host in {:.2} ms",
1620                if ok { "" } else { " (failed → fresh)" },
1621                t0.elapsed().as_secs_f64() * 1e3
1622            );
1623        }
1624        ok
1625    }
1626
1627    fn apply_kv_reuse_plan(
1628        &mut self,
1629        reuse_from: usize,
1630        plan: ReusePlan,
1631        layers: &[ReuseLayer],
1632    ) -> bool {
1633        let kv_id = self.graph_kv_id;
1634        match plan {
1635            ReusePlan::Fresh => return false,
1636            ReusePlan::Ready => {}
1637            ReusePlan::Pull(pulls) => {
1638                // Mirrors of one uniform geometry: one batched read serves
1639                // every layer. Per-layer geometry (MiMo-V2: 4/8 KV heads,
1640                // narrow V, sliding rings) is read layer by layer in the
1641                // host layout instead.
1642                let (nkv, hd) = {
1643                    let c = &self.kv_cache.layers[pulls[0].0];
1644                    (c.num_kv_heads, c.head_dim)
1645                };
1646                let uniform = pulls.iter().all(|&(li, _, _)| {
1647                    let c = &self.kv_cache.layers[li];
1648                    (c.num_kv_heads, c.head_dim) == (nkv, hd)
1649                });
1650                let batched = if uniform {
1651                    crate::gpu::graph_kv_read_rows(kv_id, &pulls, nkv, hd)
1652                } else {
1653                    None
1654                };
1655                let rows: Vec<(Vec<f32>, Vec<f32>)> = match batched {
1656                    Some(rows) => rows,
1657                    None => {
1658                        let mut rows = Vec::with_capacity(pulls.len());
1659                        for &(li, from, to) in &pulls {
1660                            let (lnkv, lhd) = {
1661                                let c = &self.kv_cache.layers[li];
1662                                (c.num_kv_heads, c.head_dim)
1663                            };
1664                            let Some((k, v, first_valid)) =
1665                                crate::gpu::graph_kv_pull_host(kv_id, li, from, to, lnkv, lhd)
1666                            else {
1667                                return false;
1668                            };
1669                            // The host continues at `to`: a sliding layer
1670                            // reads back only its last window, a full one
1671                            // every row it lacks.
1672                            let need_from = match self.layer_window(li) {
1673                                Some(w) => from.max((to + 1).saturating_sub(w)),
1674                                None => from,
1675                            };
1676                            if first_valid > need_from {
1677                                return false;
1678                            }
1679                            rows.push((k, v));
1680                        }
1681                        rows
1682                    }
1683                };
1684                for ((li, from, to), (k, v)) in pulls.into_iter().zip(rows) {
1685                    let cache = &mut self.kv_cache.layers[li];
1686                    let row = cache.num_kv_heads * cache.head_dim;
1687                    for p in 0..to - from {
1688                        cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
1689                    }
1690                    if cache.pos_len() != to {
1691                        return false;
1692                    }
1693                }
1694            }
1695        }
1696        // A mirror past the prefix (a greedy burst that ran beyond the stop)
1697        // holds rows of the OLD continuation: rewind it so the next graph
1698        // token re-syncs those positions from the host.
1699        for (li, l) in layers.iter().enumerate() {
1700            if l.full
1701                && l.device_rows.is_some_and(|d| d > reuse_from)
1702                && !crate::gpu::graph_kv_set_stored(kv_id, li, reuse_from)
1703            {
1704                return false;
1705            }
1706        }
1707        true
1708    }
1709
1710    /// Finish a generation lifecycle after the MTP/router owners were
1711    /// detached.  Every terminal path must put those owners back before the
1712    /// pooled pipeline can serve another request.  Graph side channels and
1713    /// device mirrors are cleared on errors and cancellations; a successful
1714    /// generation keeps its decode-ready host cache for KV reuse.
1715    fn finish_generation(
1716        &mut self,
1717        mtp: &mut Option<MtpModule>,
1718        router: &mut Option<crate::swarm::DynRouter>,
1719        clear_sequence: bool,
1720    ) {
1721        // A dynamic route may have switched the overlay before the terminal
1722        // path. Restore the backbone while the detached router is still
1723        // available, because set_active_skill also owns the overlay reset.
1724        if router.is_some() {
1725            let _ = self.set_active_skill(None);
1726        }
1727        // The last speculative round's replay may still be in flight on
1728        // the second queue: whoever reads the host cache after generate()
1729        // returns (session export, the network split's KV wire, a KV
1730        // reuse) must see the final states.
1731        // A replay that failed leaves the GDN owners half-written: fail
1732        // closed and drop the sequence instead of handing the cache on.
1733        #[cfg(target_os = "macos")]
1734        let clear_sequence = clear_sequence || !crate::gpu_metal::wait_replay();
1735        if clear_sequence {
1736            self.clear_sequence_state();
1737            if let Some(m) = mtp.as_mut() {
1738                // The MTP owner is detached while generation runs, so the
1739                // trunk reset above cannot clear its host cache.  Drop its
1740                // partial rows before reattaching it to the pooled pipeline;
1741                // the next request must start from the same empty anchor on
1742                // CPU and on the device mirror.
1743                m.kv.clear();
1744            }
1745            if let Some(m) = self.mtp.as_mut() {
1746                // A non-speculative request leaves the configured MTP owner
1747                // attached.  Clear that dormant cache too when a shared
1748                // generation failure/cancellation resets the sequence.
1749                m.kv.clear();
1750            }
1751        }
1752        self.graph_want_logits = false;
1753        self.graph_head_required = false;
1754        self.graph_logits = None;
1755        self.graph_failed
1756            .store(false, std::sync::atomic::Ordering::Relaxed);
1757        self.cancel
1758            .store(false, std::sync::atomic::Ordering::Relaxed);
1759        self.dyn_router = router.take().or(self.dyn_router.take());
1760        self.mtp = mtp.take().or(self.mtp.take());
1761        self.mtp_graph_mode = None;
1762        self.spec_forced = None;
1763    }
1764
1765    /// Consume a graph failure reported by a forward that returns only a
1766    /// hidden vector.  `forward_ids` is a public Result API, so it must not
1767    /// turn the graph's zero hidden sentinel into a valid lm_head result.
1768    fn check_forward_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1769        if self
1770            .graph_failed
1771            .swap(false, std::sync::atomic::Ordering::Relaxed)
1772        {
1773            self.cancel
1774                .store(false, std::sync::atomic::Ordering::Relaxed);
1775            self.clear_sequence_state();
1776            self.graph_logits = None;
1777            self.graph_want_logits = false;
1778            self.graph_head_required = false;
1779            return Err(format!("GPU graph failed during {phase} at position {pos}"));
1780        }
1781        Ok(())
1782    }
1783
1784    #[cfg(target_os = "macos")]
1785    fn fail_metal_graph(&mut self, reason: &str) {
1786        crate::pipeline::METAL_GRAPH_ERRORS
1787            .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
1788        self.clear_sequence_state();
1789        self.graph_logits = None;
1790        self.graph_failed
1791            .store(true, std::sync::atomic::Ordering::Relaxed);
1792        self.cancel
1793            .store(true, std::sync::atomic::Ordering::Relaxed);
1794        tracing::error!("native Metal TokenGraph failed closed: {reason}");
1795    }
1796
1797    /// Start an NLL/PPL request with all graph side channels in a known
1798    /// state.  A graph failure also raises the cooperative cancel bit; it is
1799    /// consumed here and that graph-induced bit is cleared so an independent
1800    /// request can be reused.  A caller-owned cancellation remains intact.
1801    fn nll_begin(&mut self) -> Result<(), String> {
1802        if self
1803            .graph_failed
1804            .swap(false, std::sync::atomic::Ordering::Relaxed)
1805        {
1806            self.cancel
1807                .store(false, std::sync::atomic::Ordering::Relaxed);
1808            self.clear_sequence_state();
1809            self.graph_logits = None;
1810            self.graph_want_logits = false;
1811            self.graph_head_required = false;
1812            return Err("GPU graph failed before NLL scoring".to_string());
1813        }
1814        self.clear_sequence_state();
1815        self.graph_logits = None;
1816        self.graph_want_logits = false;
1817        self.graph_head_required = false;
1818        Ok(())
1819    }
1820
1821    /// End an NLL/PPL request, including the side channels that are not part
1822    /// of the host KV cache.  This is intentionally explicit instead of
1823    /// relying on a tuple/sentinel return: callers must see every failure.
1824    fn nll_end(&mut self) {
1825        self.clear_sequence_state();
1826        self.graph_logits = None;
1827        self.graph_want_logits = false;
1828        self.graph_head_required = false;
1829        self.graph_failed
1830            .store(false, std::sync::atomic::Ordering::Relaxed);
1831    }
1832
1833    /// Check the graph failure channel at a scoring boundary and leave the
1834    /// pipeline reusable when the device path failed.
1835    fn nll_check_graph(&mut self, phase: &str, pos: usize) -> Result<(), String> {
1836        #[cfg(test)]
1837        if self.nll_test_fail_at == Some(pos) {
1838            self.nll_test_fail_at = None;
1839            self.graph_failed
1840                .store(true, std::sync::atomic::Ordering::Relaxed);
1841            self.cancel
1842                .store(true, std::sync::atomic::Ordering::Relaxed);
1843        }
1844        if self
1845            .graph_failed
1846            .swap(false, std::sync::atomic::Ordering::Relaxed)
1847        {
1848            self.cancel
1849                .store(false, std::sync::atomic::Ordering::Relaxed);
1850            self.clear_sequence_state();
1851            self.graph_logits = None;
1852            self.graph_want_logits = false;
1853            return Err(format!(
1854                "GPU graph failed during NLL {phase} at position {pos}"
1855            ));
1856        }
1857        Ok(())
1858    }
1859
1860    /// Map a virtual layer index to its physical weight index.
1861    /// Looped Transformer (Nanbeige 4.2): 22 physical layers × 2 loops = 44 virtual;
1862    /// virtual layer 23 maps back to physical layer 1 (23 % 22 = 1).
1863    #[inline]
1864    pub fn phys_layer(&self, virtual_idx: usize) -> usize {
1865        virtual_idx % self.physical_layers
1866    }
1867
1868    /// True when `virtual_idx` is the last layer of a loop iteration
1869    /// (used for loop_final_norm insertion).
1870    #[inline]
1871    pub fn is_loop_end(&self, virtual_idx: usize) -> bool {
1872        self.loop_final_norm && (virtual_idx + 1) % self.physical_layers == 0
1873    }
1874
1875    /// Build a pipeline from parts (used by the loader and tests).
1876    #[allow(clippy::too_many_arguments)]
1877
1878    /// Whole-block q1 token graph on the GPU (macOS/Metal): the run of
1879    /// consecutive q1 layers — GDN *and* full attention — starting at
1880    /// `start` executes as few command buffers as the CPU truly needs.
1881    /// Hidden stays device-resident across every layer; the only syncs
1882    /// are before each CPU attend (it needs q/k/v and owns the KV
1883    /// cache) and the final hidden readback. Recurrent states
1884    /// round-trip through shared memory (the CPU stays their owner, so
1885    /// every other path remains coherent). Returns the first layer
1886    /// index NOT covered (== `start` → refused, caller falls through
1887    /// to the per-layer CPU path).
1888    /// Should prefill run position-by-position through the GPU token
1889    /// graph instead of the batched CPU chunk-GEMM? True for q1 GDN
1890    /// hybrids on native Metal: their chunk prefill is walled by the
1891    /// sequential scalar recurrence, so the graph's decode rate wins.
1892    /// NOT for Looped Transformers, despite the per-chunk loop_final_norm
1893    /// sync: the chunk-GEMM amortizes each weight over the whole chunk,
1894    /// which the per-position graph cannot (Nanbeige 4.2 on M4, 512-token
1895    /// prompt: 85 tok/s chunked vs 14 through the graph).
1896    #[cfg(target_os = "macos")]
1897    fn graph_prefill_preferred(&self) -> bool {
1898        let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
1899        if !crate::gpu::enabled_here()
1900            || !graph_force
1901            || std::env::var("CMF_GPU_BLOCK")
1902                .map(|v| v == "0")
1903                .unwrap_or(false)
1904            // CMF_PREFILL_GRAPH=0: the chunked prefill (GEMM projections,
1905            // CPU recurrence) instead of the per-position token graph.
1906            || std::env::var("CMF_PREFILL_GRAPH").as_deref() == Ok("0")
1907        {
1908            return false;
1909        }
1910        self.weights
1911            .layers
1912            .iter()
1913            .any(|lw| {
1914                matches!(&lw.attn, AttnKind::LinearGdn(w) if w.in_proj_qkv.metal_graph_parts().is_some())
1915            })
1916    }
1917
1918    /// Prompt ingest through the batched wgpu graph in device-prefix mode:
1919    /// a MoE stack that does not fit the card runs each chunk's leading
1920    /// layers on the device (experts resident) and the rest on the host's
1921    /// batched walk. On by default for models with per-layer attention
1922    /// geometry (MiMo-V2 — its measured default); `CMF_BATCH_PREFIX=1`
1923    /// opts any other MoE model in, `=0` keeps the chunked host prefill.
1924    #[cfg(not(target_os = "macos"))]
1925    fn batch_prefix_prefill(&self) -> bool {
1926        let forced = match std::env::var("CMF_BATCH_PREFIX").as_deref() {
1927            Ok("0") => return false,
1928            Ok("1") => true,
1929            _ => false,
1930        };
1931        crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
1932            && crate::gpu::enabled_here()
1933            && !self.graph_refused()
1934            && (forced || self.graph_attn_decline_reason().is_some())
1935            && self.wgpu_graph_attn_decline().is_none()
1936            && self.attn_softcap == 0.0
1937            && self
1938                .weights
1939                .layers
1940                .iter()
1941                .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
1942            && self.automatic_gpu_prefix().is_some()
1943    }
1944
1945    #[cfg(not(target_os = "macos"))]
1946    fn graph_prefill_preferred(&self) -> bool {
1947        // Discrete-GPU wgpu whole-token graph: GDN layers carry recurrent state
1948        // (conv ring + delta-rule S) resident on the GPU. A batched CPU prefill
1949        // builds that state on the CPU only, leaving the GPU buffers zeroed at
1950        // decode → garbage. Route GDN-hybrid prefill through the graph one
1951        // position at a time so the resident state is seeded exactly as decode
1952        // will read it. Pure-attention models keep the batched CPU prefill (its
1953        // KV mirror re-syncs from the CPU cache, so no seeding gap).
1954        let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill);
1955        if !graph_on || !crate::gpu::enabled_here() {
1956            return false;
1957        }
1958        // Embryo's phase state is device-owned by the resident graph during
1959        // prefill; the batched CPU path would leave decode seeing a zeroed
1960        // device recurrence.  Route the prompt position-by-position too.
1961        if self.embryo_resident_eligible() {
1962            // The resident graph owns the phase state and anchor KV on the
1963            // device.  A prefill-only graph would leave decode on the host
1964            // with no way to import that state, so keep the whole sequence
1965            // on one owner (or use the ordinary CPU prefill/decode pair).
1966            return crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
1967        }
1968        // The descriptor-aware Prism graph now carries both the FWHT/affine
1969        // transforms and resident GDN state, so it is also the exact prefill
1970        // path for this model.  Keeping it here (rather than falling through
1971        // to the CPU chunk walk) is required for a long prompt to seed the
1972        // same device state that decode consumes.
1973        // O(1) needs the CPU prefill: the q-trace that seals the Nyström
1974        // skeleton is recorded there and nowhere else. The GDN half of
1975        // the hybrid loses nothing — the graph's first decode creates
1976        // its (ring, S) entries seeded from `cpu_state`, the same
1977        // handoff every graph run relies on when the entry is fresh.
1978        // Without this line the two designs collide on hybrids and o1
1979        // never becomes graph-portable: prefill through the graph
1980        // records no trace, so views stay None forever.
1981        if self.o1_active() {
1982            return false;
1983        }
1984        // A model the wgpu graphs decline outright would walk its prompt
1985        // one position at a time through a graph that never runs: the
1986        // batched CPU chunk prefill is the right ingest for it.
1987        if self.wgpu_graph_attn_decline().is_some() {
1988            return false;
1989        }
1990        if self
1991            .weights
1992            .layers
1993            .iter()
1994            .any(|lw| matches!(&lw.attn, AttnKind::LinearGdn(_)))
1995        {
1996            return true;
1997        }
1998        // MoE models too: the chunked CPU prefill runs every expert on the
1999        // host (Hy-MT2-30B-A3B on a Xeon: 8 tok/s of ingest against 53 of
2000        // graph decode), while the token graph — and the batched graph under
2001        // CMF_BATCH_K — keep the experts resident. Full attention in the
2002        // graph writes the KV mirror that decode reads, exactly as it does
2003        // for the hybrids' attention layers. Only when the whole stack is
2004        // resident: with a device prefix the per-position walk finishes
2005        // every token on the host, and the chunked prefill (GEMMs on the
2006        // card, the expert loop batched on the host) is the faster ingest
2007        // (the 8 GB ladder point: 7 tok/s chunked against ~1 walked).
2008        self.weights
2009            .layers
2010            .iter()
2011            .any(|lw| matches!(&lw.ffn, FfnKind::Moe(_)))
2012            && self.automatic_gpu_prefix().is_none()
2013    }
2014
2015    #[cfg(target_os = "macos")]
2016    fn q1_graph_gpu(
2017        &mut self,
2018        start: usize,
2019        upto: Option<usize>,
2020        position: usize,
2021        h: &mut [f32],
2022    ) -> usize {
2023        let _mt0 = std::time::Instant::now(); // CMF_METAL_HOSTPROF
2024        use crate::gpu::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, GraphDims, MetalFfn, TokenGraph};
2025        let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
2026        if self.attn_softcap > 0.0 // capped scores: no graph kernel — CPU path
2027            || !crate::gpu::enabled_here()
2028            || !graph_force
2029            || std::env::var("CMF_GPU_BLOCK")
2030                .map(|v| v == "0")
2031                .unwrap_or(false)
2032        {
2033            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2034                eprintln!(
2035                    "block-graph: front gate (softcap={} enabled_here={} graph_force={})",
2036                    self.attn_softcap > 0.0,
2037                    crate::gpu::enabled_here(),
2038                    graph_force,
2039                );
2040            }
2041            if self.graph_head_required {
2042                self.fail_metal_graph("native graph front gate refused");
2043            }
2044            return start;
2045        }
2046        // The graph encodes SiLU or exact-GELU FFNs and attention with an
2047        // explicit model scale. Sliding-window layers with their own RoPE
2048        // table / rotary width ride it too (`metal_graph_swa`): both the
2049        // device attend and the sandwich's CPU attend read each layer's
2050        // window and table. Sandwich norms and other activations fall back
2051        // to the CPU path.
2052        let swa_graph = self.metal_graph_swa();
2053        if (self.swa.is_some() && !swa_graph)
2054            || self.global_attn.is_some()
2055            || self.attention_heads_per_layer.is_some()
2056            || self.attn_v_norm
2057            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
2058            || (self.graph_attn_decline_reason().is_some() && !swa_graph)
2059            || self.weights.layers.iter().any(|lw| {
2060                lw.attn_out_norm.is_some()
2061                    || lw.ffn_out_norm.is_some()
2062                    || lw.layer_scale.is_some()
2063                    || matches!(&lw.ffn, FfnKind::Dense(d) if !matches!(d.act, Act::Silu | Act::Gelu))
2064            })
2065        {
2066            // The Metal graphs' device attend has no per-layer attention
2067            // geometry (the wgpu graphs do): say so once, by name.
2068            if let Some(reason) = self.graph_attn_decline_reason().filter(|_| !swa_graph) {
2069                self.note_graph_decline("metal block graph", reason);
2070            }
2071            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2072                eprintln!(
2073                    "block-graph: arch ineligible (swa={} gattn={} hpl={} vnorm={} scale_delta={:.2e})",
2074                    self.swa.is_some(),
2075                    self.global_attn.is_some(),
2076                    self.attention_heads_per_layer.is_some(),
2077                    self.attn_v_norm,
2078                    (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs(),
2079                );
2080            }
2081            if self.graph_head_required {
2082                self.fail_metal_graph("native graph architecture gate refused");
2083            }
2084            return start;
2085        }
2086        // Looped Transformer: the graph covers ALL loop iterations;
2087        // encode_loop_norm is inserted on-device at each boundary.
2088        let limit = upto
2089            .map(|u| u + 1)
2090            .unwrap_or(self.num_layers)
2091            .min(self.num_layers);
2092
2093        enum Item<'a> {
2094            Gdn {
2095                run: Vec<GdnGpuLayer<'a>>,
2096                first: usize,
2097            },
2098            Attn {
2099                l: AttnGpuLayer<'a>,
2100                li: usize,
2101                q_norm: Option<&'a [f32]>,
2102                k_norm: Option<&'a [f32]>,
2103                output_gate: bool,
2104                /// `self_attn.g_proj` output gate (per-head flag): the
2105                /// sandwich projects the device-normed input on the host
2106                /// and gates the attend's output before O.
2107                proj_gate: Option<(&'a QTensor, bool)>,
2108                /// The same gate's f32 rows `[nh × hidden]` when the device
2109                /// attend can apply it (per-head sigmoid, f32 in RAM).
2110                head_gate_w: Option<&'a [f32]>,
2111                bias: Option<(&'a [f32], &'a [f32], &'a [f32])>,
2112                /// Attend on the device too (no sync): F32 KV, no
2113                /// o1/bias, dims inside the kernels' contract.
2114                full_gpu: bool,
2115            },
2116        }
2117
2118        // Device-attend KERNEL contract, shared by every Full layer. The
2119        // hd>128 default-off POLICY is applied after the scan: it was
2120        // measured on dense models, and a MoE plan inverts it — with the
2121        // experts on device each CPU-attend sandwich costs a
2122        // commit+wait, ~30 submits/token (W2 on M4: 14.7 tok/s
2123        // sandwiched vs 27.1 device-attend vs 18.8 pure CPU).
2124        let attend_mode = std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into());
2125        let attend_contract = attend_mode != "0"
2126            && attend_mode != "off"
2127            && self.head_dim % 4 == 0
2128            && self.head_dim <= 256
2129            && self.rotary_dim >= 2
2130            && self.rotary_dim <= self.head_dim
2131            && (self.rotary_dim / 2) % 32 == 0
2132            && self.num_kv_heads > 0
2133            && self.num_heads % self.num_kv_heads == 0;
2134
2135        let mut plan: Vec<Item> = Vec::new();
2136        let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
2137        // Break-reason diagnostics ride the same env as the plan summary.
2138        let block_diag = std::env::var("CMF_GRAPH_DBG").is_ok();
2139        let mut scan = start;
2140        while scan < limit {
2141            let lw = &self.weights.layers[self.phys_layer(scan)];
2142            let ffn = match &lw.ffn {
2143                FfnKind::Dense(d) if d.segs.is_empty() => {
2144                    let (Some(g), Some(u), Some(dn)) = (
2145                        d.gate_proj.metal_graph_parts(),
2146                        d.up_proj.metal_graph_parts(),
2147                        d.down_proj.metal_graph_parts(),
2148                    ) else {
2149                        if block_diag {
2150                            eprintln!(
2151                                "block-graph: L{scan} FFN trio not graph-mappable — run ends"
2152                            );
2153                        }
2154                        break;
2155                    };
2156                    MetalFfn::Dense {
2157                        gate: g,
2158                        up: u,
2159                        down: dn,
2160                        // The front gate admitted SiLU or exact GELU only.
2161                        gelu: d.act == Act::Gelu,
2162                    }
2163                }
2164                FfnKind::Moe(m) => {
2165                    let Some(moe) = metal_moe_graph_parts(m, self.hidden_size) else {
2166                        if block_diag {
2167                            eprintln!(
2168                                "block-graph: L{scan} MoE outside the graph contract — run ends"
2169                            );
2170                        }
2171                        break;
2172                    };
2173                    if let QTensor::Mapped { model, .. } = &m.experts[0].gate_proj {
2174                        model_ref.get_or_insert_with(|| model.clone());
2175                    }
2176                    MetalFfn::Moe(moe)
2177                }
2178                _ => {
2179                    if block_diag {
2180                        eprintln!("block-graph: L{scan} non-graph FFN — run ends");
2181                    }
2182                    break;
2183                }
2184            };
2185            match &lw.attn {
2186                AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
2187                    let parts = (
2188                        w.in_proj_qkv.metal_graph_parts(),
2189                        w.in_proj_z.metal_graph_parts(),
2190                        w.in_proj_a.f32_parts(),
2191                        w.in_proj_b.f32_parts(),
2192                        w.out_proj.metal_graph_parts(),
2193                    );
2194                    let (Some(qkv), Some(z), Some(a), Some(b), Some(out)) = parts else {
2195                        if block_diag {
2196                            eprintln!(
2197                                "block-graph: L{scan} GDN parts refused (qkv={} z={} a_f32={} b_f32={} out={})",
2198                                w.in_proj_qkv.metal_graph_parts().is_some(),
2199                                w.in_proj_z.metal_graph_parts().is_some(),
2200                                w.in_proj_a.f32_parts().is_some(),
2201                                w.in_proj_b.f32_parts().is_some(),
2202                                w.out_proj.metal_graph_parts().is_some(),
2203                            );
2204                        }
2205                        break;
2206                    };
2207                    if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
2208                        model_ref.get_or_insert_with(|| model.clone());
2209                    }
2210                    let gl = GdnGpuLayer {
2211                        attn_norm: &lw.input_norm,
2212                        post_norm: &lw.post_norm,
2213                        qkv,
2214                        z,
2215                        a,
2216                        b,
2217                        out,
2218                        ffn,
2219                        conv1d: &w.conv1d,
2220                        a_log: &w.a_log,
2221                        dt_bias: &w.dt_bias,
2222                        gnorm: &w.norm,
2223                    };
2224                    match plan.last_mut() {
2225                        Some(Item::Gdn { run, .. }) => run.push(gl),
2226                        _ => plan.push(Item::Gdn {
2227                            run: vec![gl],
2228                            first: scan,
2229                        }),
2230                    }
2231                }
2232                AttnKind::Full {
2233                    wq,
2234                    wk,
2235                    wv,
2236                    wo,
2237                    q_norm,
2238                    k_norm,
2239                    output_gate,
2240                    softplus_gate,
2241                    bias,
2242                } if (!self.kv_cache.layers[scan].o1_sealed()
2243                    // Sealed o1 stays plannable when the Metal o1 port
2244                    // is on: full_gpu attends through the device state,
2245                    // and any refusal falls to the sandwich, whose CPU
2246                    // core routes sealed layers through the nystrom step.
2247                    || std::env::var("CMF_O1_METAL").as_deref() == Ok("1"))
2248                    // A projected output gate: Spark-X2.5's per-head
2249                    // sigmoid form only (Laguna's softplus gate is
2250                    // unmeasured on these paths and keeps the CPU walk).
2251                    && softplus_gate
2252                        .as_ref()
2253                        .is_none_or(|(_, per_head)| *per_head && self.proj_gate_sigmoid) =>
2254                {
2255                    let parts = (
2256                        wq.metal_graph_parts(),
2257                        wk.metal_graph_parts(),
2258                        wv.metal_graph_parts(),
2259                        wo.metal_graph_parts(),
2260                    );
2261                    let (Some(pq), Some(pk), Some(pv), Some(po)) = parts else {
2262                        break;
2263                    };
2264                    if let QTensor::Mapped { model, .. } = wq {
2265                        model_ref.get_or_insert_with(|| model.clone());
2266                    }
2267                    let cache = &self.kv_cache.layers[scan];
2268                    // O(1) layer on Metal: the device attends through the
2269                    // sealed Nystrom state (opt-in while the port proves
2270                    // itself). Unsealed -> sandwich path = the CPU o1 step.
2271                    let o1_metal = cache.o1.is_some()
2272                        && std::env::var("CMF_O1_METAL").as_deref() == Ok("1")
2273                        && cache.o1_views().is_some();
2274                    // The device attend takes the layer's window and RoPE
2275                    // table, and the per-head gate from f32 rows; a gate
2276                    // held otherwise sandwiches.
2277                    let head_gate_w = softplus_gate
2278                        .as_ref()
2279                        .and_then(|(g, _)| g.f32_parts())
2280                        .filter(|&(_, r, c)| r == self.num_heads && c == self.hidden_size)
2281                        .map(|(d, _, _)| d);
2282                    let full_gpu = attend_contract
2283                        && softplus_gate.is_none() == head_gate_w.is_none()
2284                        && cache.mode == crate::kv_cache::KvMode::F32
2285                        && (cache.o1.is_none() || o1_metal)
2286                        && bias.is_none()
2287                        && pq.1 == self.num_heads * self.head_dim * (1 + *output_gate as usize)
2288                        && pk.1 == self.num_kv_heads * self.head_dim
2289                        && pv.1 == self.num_kv_heads * self.head_dim
2290                        && po.2 == self.num_heads * self.head_dim;
2291                    plan.push(Item::Attn {
2292                        l: AttnGpuLayer {
2293                            attn_norm: &lw.input_norm,
2294                            post_norm: &lw.post_norm,
2295                            wq: pq,
2296                            wk: pk,
2297                            wv: pv,
2298                            wo: po,
2299                            ffn,
2300                        },
2301                        li: scan,
2302                        q_norm: q_norm.as_deref(),
2303                        k_norm: k_norm.as_deref(),
2304                        output_gate: *output_gate,
2305                        proj_gate: softplus_gate.as_ref().map(|(g, per_head)| (g, *per_head)),
2306                        head_gate_w,
2307                        bias: bias
2308                            .as_ref()
2309                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
2310                        full_gpu,
2311                    });
2312                }
2313                _ => break,
2314            }
2315            scan += 1;
2316        }
2317        let Some(model) = model_ref else {
2318            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2319                eprintln!("q1-graph: no model ref (start {start}, scanned to {scan})");
2320            }
2321            if self.graph_head_required {
2322                self.fail_metal_graph("native graph has no mapped model reference");
2323            }
2324            return start;
2325        };
2326        if plan.is_empty() {
2327            if std::env::var("CMF_GRAPH_DBG").is_ok() {
2328                eprintln!("q1-graph: empty plan at layer {start}");
2329            }
2330            if self.graph_head_required {
2331                self.fail_metal_graph("native graph plan is empty");
2332            }
2333            return start;
2334        }
2335        let has_moe = plan.iter().any(|it| match it {
2336            Item::Gdn { run, .. } => run.iter().any(|l| matches!(l.ffn, MetalFfn::Moe(_))),
2337            Item::Attn { l, .. } => matches!(l.ffn, MetalFfn::Moe(_)),
2338        });
2339        let has_gdn = plan.iter().any(|it| matches!(it, Item::Gdn { .. }));
2340        let dev_attend = attend_contract
2341            && (self.head_dim <= 128
2342                || has_moe
2343                // A GDN hybrid attends on a quarter of its layers: the
2344                // hd>128 caution was measured on pure-dense models where
2345                // gqa_attend dominates, and on Qwen3.8-27B (hd 256, 48
2346                // GDN + 16 attn) the sandwich costs 2x the whole decode
2347                // (1.2 vs 2.21 tok/s measured before the arena fix).
2348                || (self.head_dim <= 256 && has_gdn)
2349                // Sliding-window layers bound most attends by the window:
2350                // Spark-X2.5-1.7B q4tp on the M4 decodes 67 tok/s
2351                // device-attended against 23 sandwiched (28 syncs/token).
2352                || (self.head_dim <= 256 && swa_graph)
2353                || attend_mode == "force"
2354                || attend_mode == "256");
2355        if !dev_attend {
2356            for it in &mut plan {
2357                if let Item::Attn { li, full_gpu, .. } = it {
2358                    // The hd>128 policy is about gqa_attend; an o1 layer
2359                    // attends through its own kernel set.
2360                    let keep_o1 = self.kv_cache.layers[*li].o1.is_some()
2361                        && std::env::var("CMF_O1_METAL").as_deref() == Ok("1");
2362                    if !keep_o1 {
2363                        *full_gpu = false;
2364                    }
2365                }
2366            }
2367        }
2368        if std::env::var("CMF_GRAPH_DBG").is_ok() {
2369            use std::sync::atomic::{AtomicBool, Ordering};
2370            static SAID: AtomicBool = AtomicBool::new(false);
2371            if !SAID.swap(true, Ordering::Relaxed) {
2372                let fg = plan
2373                    .iter()
2374                    .filter(|it| matches!(it, Item::Attn { full_gpu: true, .. }))
2375                    .count();
2376                let att = plan
2377                    .iter()
2378                    .filter(|it| matches!(it, Item::Attn { .. }))
2379                    .count();
2380                eprintln!(
2381                    "q1-graph: plan of {} items from layer {start} to {scan} | dev_attend={dev_attend} full_gpu {fg}/{att} | hd={} rd={} nkv={} nh={}",
2382                    plan.len(),
2383                    self.head_dim,
2384                    self.rotary_dim,
2385                    self.num_kv_heads,
2386                    self.num_heads,
2387                );
2388            }
2389        }
2390        let dims = GraphDims {
2391            hidden: self.hidden_size,
2392            eps: self.rms_eps as f32,
2393            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2394        };
2395        let Some(mut graph) = TokenGraph::new(&model, dims, h) else {
2396            if self.graph_head_required {
2397                self.fail_metal_graph("native TokenGraph allocation refused");
2398            }
2399            return start;
2400        };
2401        if swa_graph && self.head_dim > 128 {
2402            // Spark-X2.5 (hd 256, 4 Q heads per KV head): the GQA-shared
2403            // split-K attend reads each K/V row once for the group instead
2404            // of once per head plus the importance re-read — at depth 1000
2405            // on the M4, 59.5 tok/s against 50.0 with the per-head kernel
2406            // up to the default 512 (every sliding layer sits at <= 512).
2407            graph.set_attend_blk_from(64);
2408        }
2409        let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
2410            nv: cfg.num_v_heads,
2411            nk: cfg.num_k_heads,
2412            dk: cfg.key_head_dim,
2413            dv: cfg.value_head_dim,
2414            kk: cfg.conv_kernel,
2415            hidden: self.hidden_size,
2416            inter: self.intermediate_size,
2417            c_dim: cfg.conv_dim(),
2418            eps: cfg.rms_eps as f32,
2419            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
2420        });
2421        // Validate the whole plan BEFORE encoding anything: after the
2422        // first sync a refused layer would leave the token
2423        // half-executed, so truncate to the provably encodable prefix.
2424        let mut valid = 0usize;
2425        let mut end = start;
2426        crate::gpu::stageprof(1, _mt0.elapsed()); // конец планирования
2427        if std::env::var("CMF_PLAN_DUMP").is_ok() {
2428            static ONCE: std::sync::Once = std::sync::Once::new();
2429            ONCE.call_once(|| {
2430                for it in &plan {
2431                    match it {
2432                        Item::Gdn { first, run } => {
2433                            eprintln!("plan: Gdn first={first} len={}", run.len())
2434                        }
2435                        Item::Attn { li, full_gpu, .. } => {
2436                            eprintln!("plan: Attn li={li} full_gpu={full_gpu}")
2437                        }
2438                    }
2439                }
2440            });
2441        }
2442        for item in &plan {
2443            let ok = match item {
2444                Item::Gdn { run, .. } => gcfg
2445                    .as_ref()
2446                    .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
2447                    .unwrap_or(false),
2448                Item::Attn { l, .. } => graph.attn_ok(l),
2449            };
2450            if !ok {
2451                if block_diag {
2452                    eprintln!(
2453                        "block-graph: plan item {} ({}) failed graph preflight",
2454                        valid,
2455                        match item {
2456                            Item::Gdn { run, first } => format!("GDN run L{first}+{}", run.len()),
2457                            Item::Attn { li, .. } => format!("Attn L{li}"),
2458                        }
2459                    );
2460                }
2461                break;
2462            }
2463            valid += 1;
2464            end += match item {
2465                Item::Gdn { run, .. } => run.len(),
2466                Item::Attn { .. } => 1,
2467            };
2468        }
2469        plan.truncate(valid);
2470        if plan.is_empty() {
2471            if self.graph_head_required {
2472                self.fail_metal_graph("native graph preflight produced no valid items");
2473            }
2474            return start;
2475        }
2476
2477        if self.graph_head_required && (upto.is_some() || end != self.num_layers) {
2478            self.fail_metal_graph("fused-head NLL requires a complete 64-layer graph");
2479            return start;
2480        }
2481
2482        // Plain dense decode (every item a device-attended full-attention
2483        // layer with a dense FFN, no O(1) state): the only plan shape the
2484        // masked-nibble q4tp matvec and the concurrent layer encoder were
2485        // measured on (MiniCPM5-2B, Qwen3-0.6B on the M4). Hybrids, MoE and
2486        // o1 layers keep the historical serial path bit for bit.
2487        // Every projection must be ONE dispatch (q1t adds an overlay pass,
2488        // Prism q2tp a transform pass — dependent pairs a concurrent
2489        // encoder would race).
2490        let one_pass = |t: (usize, usize, usize)| {
2491            use cortiq_core::TensorDtype as D;
2492            matches!(
2493                model.tensors[t.0].dtype,
2494                D::Q4TiledP | D::Q4Tiled | D::Q4Block | D::Q8Row | D::Q8_2f | D::Q1
2495            )
2496        };
2497        let dense_fast = plan.iter().all(|it| match it {
2498            Item::Attn {
2499                l, li, full_gpu, ..
2500            } => {
2501                *full_gpu
2502                    && self.kv_cache.layers[*li].o1.is_none()
2503                    && [l.wq, l.wk, l.wv, l.wo].into_iter().all(one_pass)
2504                    && match l.ffn {
2505                        MetalFfn::Dense { gate, up, down, .. } => {
2506                            one_pass(gate) && one_pass(up) && one_pass(down)
2507                        }
2508                        _ => false,
2509                    }
2510            }
2511            Item::Gdn { .. } => false,
2512        });
2513        let ab = crate::gpu_metal::dense_ab_arm().filter(|_| dense_fast);
2514        let _mv_fast = match ab {
2515            Some((bits, _)) => {
2516                graph.set_dense_concurrent_raw(bits & crate::gpu_metal::DENSE_CONC != 0);
2517                crate::gpu_metal::MvFastGuard::set_raw(bits)
2518            }
2519            None => {
2520                graph.set_dense_concurrent(dense_fast);
2521                // Spark-X2.5 q8_2f: the four-row q8_2f matvec decodes the
2522                // 1.7B at 40.5 tok/s on the M4 against 27.7 one-row.
2523                let q8r4 = if swa_graph {
2524                    crate::gpu_metal::DENSE_Q8R4
2525                } else {
2526                    0
2527                };
2528                crate::gpu_metal::MvFastGuard::set_bits(if dense_fast {
2529                    crate::gpu_metal::DENSE_MV | crate::gpu_metal::DENSE_FUSE | q8r4
2530                } else {
2531                    0
2532                })
2533            }
2534        };
2535
2536        // RoPE table and rotary width are per layer (`layer_inv_freq`,
2537        // `layer_geom`) inside the loop below.
2538        let pool = self.pool.clone();
2539        let (nh, nkv, hd, hs, eps) = (
2540            self.num_heads,
2541            self.num_kv_heads,
2542            self.head_dim,
2543            self.hidden_size,
2544            self.rms_eps,
2545        );
2546        let norm_style = self.norm_style;
2547        let gemma = norm_style == cortiq_core::NormStyle::Gemma;
2548        let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
2549        let kv_id = self.graph_kv_id;
2550        // GDN runs whose states await readback after the next sync
2551        // (device-attended layers add no sync, so several may stack).
2552        let mut pending: Vec<(usize, usize)> = Vec::new();
2553        // Device-attended layers: their K/V/imp are pulled from the
2554        // mirror after the final sync.
2555        let mut dev_attn: Vec<usize> = Vec::new();
2556        for item in &plan {
2557            let _xt0 = std::time::Instant::now();
2558            let _xkind: u32 = match item {
2559                Item::Gdn { .. } => 2,
2560                Item::Attn { .. } => 3,
2561            };
2562            // Looped Transformer: insert on-device norm at loop boundaries.
2563            if self.loop_final_norm {
2564                let item_start = match item {
2565                    Item::Gdn { first, .. } => *first,
2566                    Item::Attn { li, .. } => *li,
2567                };
2568                if item_start > start && self.is_loop_end(item_start - 1) {
2569                    graph.encode_loop_norm(&self.weights.final_norm);
2570                }
2571            }
2572            match item {
2573                Item::Gdn { run, first } => {
2574                    for l in &mut self.kv_cache.layers[*first..*first + run.len()] {
2575                        if l.linear_state.len() != want {
2576                            l.linear_state = vec![0f32; want];
2577                        }
2578                    }
2579                    let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
2580                        .iter()
2581                        .map(|l| l.linear_state.as_slice())
2582                        .collect();
2583                    let _ig = std::time::Instant::now();
2584                    if !graph.encode_gdn_run(run, &ro, gcfg.as_ref().unwrap()) {
2585                        // Unreachable: the plan was validated above.
2586                        tracing::error!("q1 graph: GDN run refused after validation");
2587                        return start;
2588                    }
2589                    // Early commit: the GPU starts the run while the
2590                    // CPU encodes the next layer (nothing to wait on).
2591                    graph.commit_kind = 2;
2592                    graph.commit();
2593                    crate::gpu::stageprof(0, _ig.elapsed());
2594                    pending.push((*first, run.len()));
2595                }
2596                Item::Attn {
2597                    l,
2598                    li,
2599                    q_norm,
2600                    k_norm,
2601                    output_gate,
2602                    proj_gate,
2603                    head_gate_w,
2604                    bias,
2605                    full_gpu,
2606                } => {
2607                    let _ia = std::time::Instant::now();
2608                    // This layer's attention geometry: its window, RoPE
2609                    // table and rotary width (Spark-X2.5 interleaves
2610                    // 512-window layers rotating all dims at θ 1e4 with
2611                    // full layers rotating a quarter at θ 5e6). Every model
2612                    // without sliding layers reads the global ones here.
2613                    let inv_freq_l = self.layer_inv_freq(*li);
2614                    let rd_l = self.layer_geom(*li).2;
2615                    let window_l = self.layer_window(*li);
2616                    let head_gate_w = *head_gate_w;
2617                    // ── Fully device-resident attention: no sync at all.
2618                    if *full_gpu {
2619                        let cache = &self.kv_cache.layers[*li];
2620                        let o1p = if cache.o1.is_some() {
2621                            match cache.o1_views() {
2622                                Some(views) => Some(crate::gpu::O1AttnParams {
2623                                    views,
2624                                    epoch: self.o1_epoch,
2625                                }),
2626                                // Sealed state gone mid-run: sandwich.
2627                                None => None,
2628                            }
2629                        } else {
2630                            None
2631                        };
2632                        let o1_layer = cache.o1.is_some();
2633                        if o1_layer && o1p.is_none() {
2634                            // fall to the sandwich (CPU o1 step)
2635                        }
2636                        let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
2637                        let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
2638                        let cpu_stored = if o1_layer { 0 } else { cpu_k[0].len() / hd };
2639                        // A trimmed tail always holds the window the
2640                        // device attend reads (`first = rows + 1 − w`).
2641                        debug_assert!(
2642                            cache.base() == 0 || window_l.is_some_and(|w| cpu_stored + 1 >= w)
2643                        );
2644                        let p = crate::gpu::AttnDeviceParams {
2645                            kv_id,
2646                            layer: *li,
2647                            nh,
2648                            nkv,
2649                            hd,
2650                            rd: rd_l,
2651                            position,
2652                            scale: self.attn_scale,
2653                            eps: eps as f32,
2654                            gemma,
2655                            late_qk_norm: self.qk_norm_after_rope,
2656                            output_gate: *output_gate,
2657                            q_norm: *q_norm,
2658                            k_norm: *k_norm,
2659                            inv_freq: &inv_freq_l,
2660                            cpu_k,
2661                            cpu_v,
2662                            cpu_stored,
2663                            cpu_gen: cache.generation(),
2664                            o1: o1p,
2665                            window: window_l,
2666                            head_gate: head_gate_w,
2667                        };
2668                        let o1_bad = o1_layer && p.o1.is_none();
2669                        if !o1_bad && graph.attn_device_ok(l, &p) && graph.encode_attn_device(l, &p)
2670                        {
2671                            // o1 layers leave no mirror row to pull.
2672                            if p.o1.is_none() {
2673                                dev_attn.push(*li);
2674                            }
2675                            graph.commit_kind = 3;
2676                            graph.commit();
2677                            // The footer below is skipped by `continue`:
2678                            // account the device-attn item here or its
2679                            // cost hides from the stage profile entirely.
2680                            crate::gpu::stageprof(_xkind, _xt0.elapsed());
2681                            continue;
2682                        }
2683                        // Mirror refused (nothing encoded) → sandwich.
2684                    }
2685                    graph.encode_attn_prefix(l);
2686                    if let Err(err) = graph.sync_checked() {
2687                        self.fail_metal_graph(&err);
2688                        return start;
2689                    }
2690                    if !pending.is_empty() {
2691                        let idxs: Vec<usize> =
2692                            pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2693                        let mut outs: Vec<&mut [f32]> = self
2694                            .kv_cache
2695                            .layers
2696                            .iter_mut()
2697                            .enumerate()
2698                            .filter(|(i, _)| idxs.binary_search(i).is_ok())
2699                            .map(|(_, s)| s.linear_state.as_mut_slice())
2700                            .collect();
2701                        graph.read_states(&mut outs);
2702                    }
2703                    let mut q_raw = attention::take_buf(l.wq.1);
2704                    let mut k = attention::take_buf(l.wk.1);
2705                    let mut v = attention::take_buf(l.wv.1);
2706                    graph.read_qkv(&mut q_raw, &mut k, &mut v);
2707                    // Projected output gate: g = G·norm(h) on the host, off
2708                    // the normed input the prefix left on the device.
2709                    let mut gate_raw = proj_gate.map(|(gp, _)| {
2710                        let mut normed = attention::take_buf(hs);
2711                        graph.read_normed(&mut normed);
2712                        let mut raw = attention::take_buf(gp.rows());
2713                        gp.matvec(&normed, &mut raw, pool.as_deref());
2714                        attention::recycle_buf(&mut normed);
2715                        raw
2716                    });
2717                    let cfg = QwenAttnCfg {
2718                        num_heads: nh,
2719                        num_kv_heads: nkv,
2720                        head_dim: hd,
2721                        hidden_size: hs,
2722                        position,
2723                        inv_freq: &inv_freq_l,
2724                        rotary_dim: rd_l,
2725                        scale: self.attn_scale,
2726                        softcap: self.attn_softcap,
2727                        window: window_l,
2728                        v_norm: false,
2729                        qk_norm_after_rope: self.qk_norm_after_rope,
2730                        gate_sigmoid: self.proj_gate_sigmoid,
2731                        q_norm: *q_norm,
2732                        k_norm: *k_norm,
2733                        output_gate: *output_gate,
2734                        softplus_gate: None,
2735                        rope_scale: 1.0,
2736                        bias: *bias,
2737                        rms_eps: eps,
2738                        norm_style,
2739                        pool: pool.as_deref(),
2740                        v_head_dim: hd,
2741                    };
2742                    // CMF_ATTN_ORACLE=1: diff the device attend against
2743                    // this CPU attend on identical inputs (bring-up).
2744                    let oracle = std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1")
2745                        || std::env::var("CMF_ATTN_DUMP").is_ok();
2746                    let _ = full_gpu;
2747                    let oracle_in = oracle.then(|| (q_raw.clone(), k.clone(), v.clone()));
2748                    let mut ao = attention::qwen_attention_core(
2749                        q_raw,
2750                        k,
2751                        v,
2752                        &mut self.kv_cache.layers[*li],
2753                        &cfg,
2754                    );
2755                    // CMF_ATTN_DUMP=<dir>: this token's rope'd Q and the layer's whole
2756                    // K/V cache as raw f32 (offline attention-statistics probes:
2757                    // block bounds, mass concentration). Needs CMF_GPU_ATTEND=0.
2758                    if let Ok(dir) = std::env::var("CMF_ATTN_DUMP") {
2759                        if let Some((qr0, k0, v0)) = oracle_in.clone() {
2760                            let (cq, _cg, _ck, _cv) =
2761                                attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2762                            let cache = &self.kv_cache.layers[*li];
2763                            let n = cache.head_keys(0).len() / hd;
2764                            let mut bytes: Vec<u8> = Vec::new();
2765                            for v in [nh as u32, nkv as u32, hd as u32, n as u32, position as u32] {
2766                                bytes.extend_from_slice(&v.to_le_bytes());
2767                            }
2768                            for v in &cq {
2769                                bytes.extend_from_slice(&v.to_le_bytes());
2770                            }
2771                            for g in 0..nkv {
2772                                for v in cache.head_keys(g) {
2773                                    bytes.extend_from_slice(&v.to_le_bytes());
2774                                }
2775                            }
2776                            for g in 0..nkv {
2777                                for v in cache.head_values(g) {
2778                                    bytes.extend_from_slice(&v.to_le_bytes());
2779                                }
2780                            }
2781                            let _ =
2782                                std::fs::write(format!("{dir}/L{li}_pos{position}.bin"), &bytes);
2783                        }
2784                    }
2785                    if let Some((qr0, k0, v0)) =
2786                        oracle_in.filter(|_| std::env::var("CMF_ATTN_ORACLE").as_deref() == Ok("1"))
2787                    {
2788                        let (cq, _cg, ck, cv) =
2789                            attention::finish_projection_debug(qr0, k0, v0, &cfg, position);
2790                        let mut h_now = vec![0f32; hs];
2791                        graph.read_h(&mut h_now);
2792                        let cache = &self.kv_cache.layers[*li];
2793                        let n_after = cache.head_keys(0).len() / hd;
2794                        // A sealed O(1) cache may have no dense current-row
2795                        // entry. The oracle is a debug probe, so let it see
2796                        // zero stored exact rows instead of underflowing.
2797                        let stored = n_after.saturating_sub(1);
2798                        let cpu_k: Vec<&[f32]> = (0..nkv)
2799                            .map(|g| &cache.head_keys(g)[..stored * hd])
2800                            .collect();
2801                        let cpu_v: Vec<&[f32]> = (0..nkv)
2802                            .map(|g| &cache.head_values(g)[..stored * hd])
2803                            .collect();
2804                        let p = crate::gpu::AttnDeviceParams {
2805                            kv_id,
2806                            layer: *li,
2807                            nh,
2808                            nkv,
2809                            hd,
2810                            // The layer's own geometry, as the CPU attend
2811                            // above used it (the projected gate is applied
2812                            // after this probe on both sides).
2813                            rd: rd_l,
2814                            position,
2815                            scale: self.attn_scale,
2816                            eps: eps as f32,
2817                            gemma,
2818                            late_qk_norm: self.qk_norm_after_rope,
2819                            output_gate: *output_gate,
2820                            q_norm: *q_norm,
2821                            k_norm: *k_norm,
2822                            inv_freq: &inv_freq_l,
2823                            cpu_k,
2824                            cpu_v,
2825                            cpu_stored: stored,
2826                            cpu_gen: cache.generation(),
2827                            o1: None,
2828                            window: window_l,
2829                            head_gate: None,
2830                        };
2831                        if let Some((dq, dk, dv, dao)) = graph.debug_attn_device(l, &p, &h_now) {
2832                            let md = |a: &[f32], b: &[f32]| {
2833                                a.iter().zip(b).fold(0f32, |m, (x, y)| m.max((x - y).abs()))
2834                            };
2835                            let nn = |a: &[f32]| a.iter().map(|x| x * x).sum::<f32>().sqrt();
2836                            eprintln!(
2837                                "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}",
2838                                nn(&cq),
2839                                md(&cq, &dq),
2840                                nn(&ck),
2841                                md(&ck, &dk),
2842                                nn(&cv),
2843                                md(&cv, &dv),
2844                                nn(&ao),
2845                                md(&ao, &dao)
2846                            );
2847                        } else {
2848                            eprintln!("attn-oracle L{li}: device probe declined");
2849                        }
2850                    }
2851                    if let (Some(raw), Some((_, per_head))) = (gate_raw.as_deref(), *proj_gate) {
2852                        // V is as wide as the head (the front gate), so ao
2853                        // is nh·hd here, as in `qwen_attention`.
2854                        attention::apply_projected_gate(
2855                            &mut ao,
2856                            raw,
2857                            per_head,
2858                            hd,
2859                            self.proj_gate_sigmoid,
2860                        );
2861                    }
2862                    if let Some(mut raw) = gate_raw.take() {
2863                        attention::recycle_buf(&mut raw);
2864                    }
2865                    graph.encode_attn_suffix(l, &ao);
2866                    // Early commit: the GPU starts O+FFN while the CPU
2867                    // encodes the following GDN run / attention prefix.
2868                    graph.commit();
2869                    attention::recycle_buf(&mut ao);
2870                }
2871            }
2872
2873            crate::gpu::stageprof(_xkind, _xt0.elapsed());
2874        }
2875        // Ride the final norm + lm_head in the same command buffer when
2876        // this run reaches the model's end and the caller wants logits:
2877        // the separate per-op lm_head submit (a full round trip) folds
2878        // into the sync that already happens here.
2879        let mut lm_rows = None;
2880        if self.graph_want_logits
2881            && upto.is_none()
2882            && end == self.num_layers
2883            && std::env::var("CMF_GPU_LMHEAD")
2884                .map(|v| v != "0")
2885                .unwrap_or(true)
2886        {
2887            if let Some(lm) = self.weights.lm_head.metal_graph_parts() {
2888                if graph.lm_head_ok(lm) {
2889                    graph.encode_lm_head(&self.weights.final_norm, lm);
2890                    lm_rows = Some(lm.1);
2891                }
2892            }
2893        }
2894        if self.graph_head_required && lm_rows.is_none() {
2895            METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2896            self.fail_metal_graph("fused graph head was requested but not encodable");
2897            return start;
2898        }
2899        let _sy0 = std::time::Instant::now();
2900        if let Err(err) = graph.sync_checked() {
2901            self.fail_metal_graph(&err);
2902            return start;
2903        }
2904        let _rs0 = std::time::Instant::now();
2905        if !pending.is_empty() {
2906            let idxs: Vec<usize> = pending.drain(..).flat_map(|(f, n)| f..f + n).collect();
2907            let mut outs: Vec<&mut [f32]> = self
2908                .kv_cache
2909                .layers
2910                .iter_mut()
2911                .enumerate()
2912                .filter(|(i, _)| idxs.binary_search(i).is_ok())
2913                .map(|(_, s)| s.linear_state.as_mut_slice())
2914                .collect();
2915            graph.read_states(&mut outs);
2916        }
2917        if std::env::var("CMF_GRAPH_HOSTPROF").as_deref() == Ok("1") {
2918            use std::sync::atomic::{AtomicU64, Ordering};
2919            static SY: AtomicU64 = AtomicU64::new(0);
2920            static RS: AtomicU64 = AtomicU64::new(0);
2921            static N: AtomicU64 = AtomicU64::new(0);
2922            SY.fetch_add((_rs0 - _sy0).as_nanos() as u64, Ordering::Relaxed);
2923            RS.fetch_add(_rs0.elapsed().as_nanos() as u64, Ordering::Relaxed);
2924            let n = N.fetch_add(1, Ordering::Relaxed) + 1;
2925            if n % 100 == 0 {
2926                eprintln!(
2927                    "postprof: sync-wait {:.1} ms/ток | read_states {:.1} ms/ток ({n})",
2928                    SY.load(Ordering::Relaxed) as f64 / n as f64 / 1e6,
2929                    RS.load(Ordering::Relaxed) as f64 / n as f64 / 1e6
2930                );
2931            }
2932        }
2933        if let Some(rows) = lm_rows {
2934            crate::gpu::hostprof_encode_done(_mt0);
2935            let mut lg = attention::take_buf(rows.min(self.vocab_size));
2936            graph.read_logits(&mut lg);
2937            crate::gpu::hostprof_total(_mt0);
2938            lg.resize(self.vocab_size, 0.0);
2939            if let Some(c) = self.final_softcap {
2940                for l in lg.iter_mut() {
2941                    *l = c * (*l / c).tanh();
2942                }
2943            }
2944            self.graph_logits = Some(lg);
2945        }
2946        graph.read_h(h);
2947        if self.graph_head_required && self.graph_logits.is_none() {
2948            METAL_GRAPH_HEAD_MISS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2949            self.fail_metal_graph("fused graph head completed without logits readback");
2950            return start;
2951        }
2952        METAL_GRAPH_TOK_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2953        METAL_GRAPH_LAYERS.fetch_add(
2954            end.saturating_sub(start) as u64,
2955            std::sync::atomic::Ordering::Relaxed,
2956        );
2957        if self.graph_head_required {
2958            METAL_GRAPH_HEAD_OK.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
2959        }
2960        // Device-attended layers: replay the CPU bookkeeping — append
2961        // the mirror's new K/V row (rope'd on the GPU) into the owner
2962        // cache, then bank this token's attention-importance mass.
2963        for li in dev_attn {
2964            let mut krow = attention::take_buf(nkv * hd);
2965            let mut vrow = attention::take_buf(nkv * hd);
2966            if crate::gpu::kv_mirror_read_last(kv_id, li, nkv, hd, &mut krow, &mut vrow) {
2967                let cache = &mut self.kv_cache.layers[li];
2968                cache.append(&krow, &vrow, &[]);
2969                let n = cache.seq_len;
2970                let mut imp = attention::take_buf(n);
2971                crate::gpu::kv_mirror_take_imp(kv_id, li, &mut imp);
2972                cache.accumulate_imp(&imp);
2973                attention::recycle_buf(&mut imp);
2974            }
2975            attention::recycle_buf(&mut krow);
2976            attention::recycle_buf(&mut vrow);
2977        }
2978        if let Some((_, arm)) = ab {
2979            crate::gpu_metal::dense_ab_record(arm, _mt0.elapsed());
2980        }
2981        end
2982    }
2983
2984    pub fn new(
2985        tokenizer: Tokenizer,
2986        weights: PipelineWeights,
2987        hidden_size: usize,
2988        intermediate_size: usize,
2989        num_heads: usize,
2990        num_kv_heads: usize,
2991        head_dim: usize,
2992        num_layers: usize,
2993        physical_layers: usize,
2994        loop_final_norm: bool,
2995        vocab_size: usize,
2996        rms_eps: f64,
2997        rope_base: f32,
2998        norm_style: NormStyle,
2999        max_seq_len: usize,
3000        sampler_config: SamplerConfig,
3001    ) -> Self {
3002        let rng = match sampler_config.seed {
3003            Some(s) => SplitMix64::new(s),
3004            None => SplitMix64::from_entropy(),
3005        };
3006        let inv_freq = std::sync::Arc::new(attention::rope_inv_freq(head_dim, rope_base));
3007        let pool = Pool::from_env();
3008        if let Some(p) = &pool {
3009            tracing::info!("worker pool: {} threads", p.n_workers());
3010            // Keep the workers on the socket that holds the weights.
3011            if let Some(model) = weights
3012                .lm_head
3013                .model_arc()
3014                .or_else(|| weights.embed_tokens.model_arc())
3015            {
3016                let regions: Vec<&[u8]> =
3017                    model.tensors.iter().map(|t| model.entry_bytes(t)).collect();
3018                p.bind_numa(&regions);
3019            }
3020        }
3021        Self {
3022            gpu_plan: None,
3023            tokenizer: std::sync::Arc::new(tokenizer),
3024            kv_cache: KvCache::new(num_layers, num_kv_heads, head_dim, max_seq_len),
3025            sampler_config,
3026            weights,
3027            hidden_size,
3028            intermediate_size,
3029            num_heads,
3030            num_kv_heads,
3031            head_dim,
3032            num_layers,
3033            physical_layers,
3034            loop_final_norm,
3035            vocab_size,
3036            rms_eps,
3037            rope_base,
3038            norm_style,
3039            rotary_dim: head_dim,
3040            attention_heads_per_layer: None,
3041            kv_heads_per_layer: None,
3042            v_head_dim: None,
3043            layer_dump: std::env::var_os("CMF_LAYER_DUMP")
3044                .filter(|v| !v.is_empty())
3045                .map(std::path::PathBuf::from),
3046            graph_declines: std::cell::RefCell::new(Vec::new()),
3047            mimo_moe: Default::default(),
3048            vmf_cfg: None,
3049            gdn_cfg: None,
3050            kda_cfg: None,
3051            g3n: None,
3052            dsv4: None,
3053            dsv41: None,
3054            dsv41_vision: None,
3055            dsv41_prefill: None,
3056            qwen4_exp: None,
3057            dsv4_mtp: Vec::new(),
3058            dspark: None,
3059            dspark_pending: Vec::new(),
3060            dspark_hist: Vec::new(),
3061            dspark_real: Vec::new(),
3062            dspark_trunk_picks: Vec::new(),
3063            dspark_exp: Vec::new(),
3064            dspark_draft_ns: 0,
3065            logit_multiplier: None,
3066            cancel: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)),
3067            graph_failed: std::sync::atomic::AtomicBool::new(false),
3068            kv_history: Vec::new(),
3069            kv_history_device: false,
3070            short_conv_cfg: None,
3071            mtp: None,
3072            mimo_mtp: None,
3073            verify_exact_moe: false,
3074            speculative: std::env::var("CMF_MTP").map(|v| v != "0").unwrap_or(true),
3075            ignore_eos: false,
3076            draft_full_streak: 0,
3077            spec_k_adapt: None,
3078            spec_acc_ewma: 0.7,
3079            rng,
3080            sampler_scratch: SamplerScratch::default(),
3081            spec_forced: None,
3082            spec_q: Vec::new(),
3083            spec_p: Vec::new(),
3084            spec_res: Vec::new(),
3085            spec_qs: Vec::new(),
3086            spec_ps: Vec::new(),
3087            spec_ress: Vec::new(),
3088            mtp_graph_mode: None,
3089            #[cfg(target_os = "macos")]
3090            metal_verify: None,
3091            inv_freq,
3092            ws: ForwardScratch::new(hidden_size),
3093            pool,
3094            model: None,
3095            dyn_force_f32: false,
3096            dyn_skill_layers: Vec::new(),
3097            dyn_active: None,
3098            dyn_blend_loaded: false,
3099            dyn_phi_layer: None,
3100            dyn_phi_ema: Vec::new(),
3101            dyn_phi_seen: 0,
3102            dyn_router: None,
3103            o1_cfg: None,
3104            o1_epoch: 0,
3105            o1_flags: Vec::new(),
3106            trace: false,
3107            calib_temp: 1.0,
3108            confidence_on: true,
3109            embed_multiplier: 1.0,
3110            attn_scale: 1.0 / (head_dim as f32).sqrt(),
3111            swa: None,
3112            swa_trim: None,
3113            sliding_layers: None,
3114            anchor_core: None,
3115            bounded_rope: None,
3116            kv_prefix: KvPrefix::default(),
3117            last_prefill_tokens: 0,
3118            inv_freq_local: None,
3119            rotary_dim_local: None,
3120            rope_scale: 1.0,
3121            rope_scale_local: 1.0,
3122            global_attn: None,
3123            inv_freq_global: None,
3124            attn_v_norm: false,
3125            qk_norm_after_rope: false,
3126            proj_gate_sigmoid: false,
3127            final_softcap: None,
3128            head_clusters: None,
3129            attn_softcap: 0.0,
3130            graph_want_logits: false,
3131            graph_head_required: false,
3132            graph_logits: None,
3133            embryo_graph: None,
3134            graph_refused: std::sync::atomic::AtomicBool::new(false),
3135            graph_kv_id: {
3136                static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
3137                NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
3138            },
3139            #[cfg(test)]
3140            nll_test_fail_at: None,
3141            #[cfg(test)]
3142            nll_test_force_serial: false,
3143        }
3144    }
3145
3146    /// Enable/disable per-layer O(1) Nyström attention. Only Full
3147    /// layers are eligible (a linear layer keeps its own operator).
3148    /// Applies to generation (`generate*`/`forward_ids`): the prompt
3149    /// pass stays exact, then the state seals after prefill or at the
3150    /// deferred skeleton-safe boundary for short prompts; decode runs on
3151    /// the O(1) state. Teacher-forced scoring (`ppl_ids`) intentionally
3152    /// stays exact.
3153    pub fn set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) {
3154        if let Err(e) = self.try_set_o1(cfg) {
3155            tracing::error!("{e}");
3156        }
3157    }
3158
3159    /// True when the file carries a native bounded anchor
3160    /// (`arch.anchor_core`): its state is a fixed record the header
3161    /// fixes, and the post-hoc O(1) override is meaningless on it.
3162    pub fn bounded_native(&self) -> bool {
3163        self.anchor_core.is_some()
3164    }
3165
3166    /// Bytes the resident device graph holds for this pipeline's sequence:
3167    /// `(recurrent state, anchor KV/ring)`; None on the host path.
3168    pub fn device_state_bytes(&self) -> Option<(u64, u64)> {
3169        crate::gpu::embryo_device_state_bytes(self.graph_kv_id)
3170    }
3171
3172    /// Why an O(1) override is refused on this pipeline, if it is.
3173    pub fn o1_refusal(&self) -> Option<String> {
3174        self.anchor_core.as_ref().map(|ac| {
3175            format!(
3176                "--o1 / CMF_O1 refused: the anchor is native bounded \
3177                 (anchor_core kind={} window={} sink={}); the file's operator \
3178                 is executed as-is and no post-hoc Nyström overlay applies",
3179                ac.kind, ac.window, ac.sink
3180            )
3181        })
3182    }
3183
3184    /// `set_o1` that reports the refusal instead of logging it.
3185    pub fn try_set_o1(&mut self, cfg: Option<crate::nystrom::O1Cfg>) -> Result<(), String> {
3186        if let Some(c) = &cfg {
3187            if let Some(why) = self.o1_refusal() {
3188                self.o1_flags = Vec::new();
3189                self.o1_cfg = None;
3190                return Err(why);
3191            }
3192            if crate::nystrom::o1_deferred_boundary(c.w, c.sink).is_none() {
3193                self.o1_flags.clear();
3194                self.o1_cfg = None;
3195                return Err(format!(
3196                    "o1 disabled: w + sink + slack + 1 overflows usize (w={}, sink={})",
3197                    c.w, c.sink
3198                ));
3199            }
3200        }
3201        self.o1_flags = match &cfg {
3202            Some(c) => {
3203                let mut flags = c.layer_flags(self.num_layers);
3204                for (li, f) in flags.iter_mut().enumerate() {
3205                    // The Nyström state replaces a full-context plain
3206                    // softmax: a sliding window or a learned sink is not
3207                    // something it can represent, and a V narrower than
3208                    // the head is not what its streaming state stores.
3209                    // Those layers keep exact cache attention.
3210                    if *f
3211                        && (!matches!(
3212                            self.weights.layers[self.phys_layer(li)].attn,
3213                            AttnKind::Full { .. }
3214                        ) || self.layer_window(li).is_some()
3215                            || self.kv_cache.layers[li].sinks.is_some()
3216                            || self.layer_v_dim(li) != self.layer_geom(li).1)
3217                    {
3218                        *f = false;
3219                    }
3220                }
3221                flags
3222            }
3223            None => Vec::new(),
3224        };
3225        if let Some(c) = &cfg {
3226            let n = self.o1_flags.iter().filter(|&&f| f).count();
3227            tracing::info!(
3228                "o1 nystrom attention: {n}/{} layer(s), m={} w={} sink={} rect={:?}",
3229                self.num_layers,
3230                c.m,
3231                c.w,
3232                c.sink,
3233                c.rect
3234            );
3235        }
3236        self.o1_cfg = cfg;
3237        Ok(())
3238    }
3239
3240    /// Install the file's bounded anchor: one fixed-size ring per
3241    /// `AttnKind::Bounded` layer (from the header, not per prompt) and
3242    /// the shared relative-rotation table. Must run after the RoPE setup
3243    /// (`set_rotary`, YaRN) so the table is built from the final
3244    /// `inv_freq`.
3245    pub fn install_bounded(
3246        &mut self,
3247        cfg: &cortiq_core::AnchorCoreConfig,
3248    ) -> Result<(), String> {
3249        if !cortiq_core::AnchorCoreConfig::KINDS.contains(&cfg.kind.as_str()) {
3250            return Err(format!(
3251                "anchor_core kind '{}' is not executable by this runtime",
3252                cfg.kind
3253            ));
3254        }
3255        if cfg.window == 0 {
3256            return Err("anchor_core.window must be >= 1".into());
3257        }
3258        let mut n = 0usize;
3259        for li in 0..self.num_layers {
3260            let pl = self.phys_layer(li);
3261            if let AttnKind::Bounded(w) = &self.weights.layers[pl].attn {
3262                if w.window != cfg.window || w.sink != cfg.sink {
3263                    return Err(format!(
3264                        "layer {li}: bounded weights (window {} sink {}) disagree with \
3265                         anchor_core (window {} sink {})",
3266                        w.window, w.sink, cfg.window, cfg.sink
3267                    ));
3268                }
3269                self.kv_cache.layers[li].install_bounded(cfg.window);
3270                n += 1;
3271            }
3272        }
3273        if n == 0 {
3274            return Err("anchor_core is present but no layer executes it".into());
3275        }
3276        let rope = crate::bounded::BoundedRope::new(cfg.window, &self.inv_freq, self.rope_scale);
3277        self.bounded_rope = Some(std::sync::Arc::new(rope));
3278        self.anchor_core = Some(cfg.clone());
3279        self.embryo_graph = None;
3280        tracing::info!(
3281            "bounded anchor {}: {n} layer(s), window {} sink {} — {} B of ring per layer",
3282            cfg.kind,
3283            cfg.window,
3284            cfg.sink,
3285            self.kv_cache.layers.iter().map(|l| l.bounded_state_bytes()).max().unwrap_or(0)
3286        );
3287        Ok(())
3288    }
3289
3290    /// Tag every layer's cache with its wire record kind and the model's
3291    /// operator identity (hash64 of `linear_core_identity` JSON) so the
3292    /// versioned state wire refuses a peer holding another operator.
3293    pub fn install_wire_identity(&mut self, identity: u64) {
3294        for li in 0..self.kv_cache.layers.len() {
3295            let pl = self.phys_layer(li);
3296            let kind = match self.weights.layers.get(pl).map(|l| &l.attn) {
3297                Some(AttnKind::Bounded(_)) => crate::kv_cache::WireKind::Bounded,
3298                Some(AttnKind::Linear(_))
3299                | Some(AttnKind::LinearGdn(_))
3300                | Some(AttnKind::ShortConv(_))
3301                | Some(AttnKind::Kda(_)) => crate::kv_cache::WireKind::Linear,
3302                _ => crate::kv_cache::WireKind::Full,
3303            };
3304            let l = &mut self.kv_cache.layers[li];
3305            l.wire_kind = kind;
3306            l.wire_identity = identity;
3307            // A per-layer geometry (`set_attn_geometry`, Gemma's global
3308            // heads) rebuilds a layer cache with index 0: re-tag it.
3309            l.wire_layer = li as u32;
3310        }
3311    }
3312
3313    /// Forget the reuse keys (legacy `kv_history` and the bounded prefix).
3314    pub fn clear_history(&mut self) {
3315        self.kv_history.clear();
3316        self.kv_history_device = false;
3317        self.kv_prefix.clear();
3318    }
3319
3320    /// Did THIS pipeline's token graph refuse for a structural reason?
3321    /// (Per pipeline: another lane's refusal, or a new pipeline of the
3322    /// same model, never changes it.)
3323    pub fn graph_refused(&self) -> bool {
3324        self.graph_refused
3325            .load(std::sync::atomic::Ordering::Relaxed)
3326    }
3327
3328    /// Remember a structural refusal of this pipeline's token graph.
3329    pub fn mark_graph_refused(&self) {
3330        if !self
3331            .graph_refused
3332            .swap(true, std::sync::atomic::Ordering::Relaxed)
3333        {
3334            tracing::info!(
3335                "token graph: unsupported for this pipeline (seq {}) — not retrying",
3336                self.graph_kv_id
3337            );
3338        }
3339    }
3340
3341    /// Position the resident Embryo graph holds for this pipeline's
3342    /// sequence (`Some(next position)`), `None` when the device holds no
3343    /// image of it (host-owned sequence, or none started).
3344    pub fn device_sequence_position(&self) -> Option<usize> {
3345        crate::gpu::embryo_device_next_position(self.graph_kv_id)
3346    }
3347
3348    /// Is the sequence this pipeline's reuse key describes owned by the
3349    /// device path it would take now? A prefix recorded on the resident
3350    /// graph continues only there (at exactly `n`), a host prefix only on
3351    /// the host; any mismatch means the next turn re-prefills from zero.
3352    fn prefix_owner_matches(&self, n: usize, recorded_on_device: bool) -> bool {
3353        let dev = self.device_sequence_position();
3354        if recorded_on_device {
3355            dev == Some(n) && self.embryo_resident_wanted()
3356        } else {
3357            dev.is_none()
3358        }
3359    }
3360
3361    /// The weights under the sequence changed (a real skill switch): every
3362    /// cached state was computed by other weights. Clear the host KV /
3363    /// ring / recurrent state AND the reuse keys — a surviving `kv_prefix`
3364    /// would let the next turn "extend" a prefix the new weights never
3365    /// saw — drop the packed resident graph (it holds the old FFN
3366    /// tensors; the next build packs the live ones under a fresh id) and
3367    /// reset its device sequence.
3368    pub(crate) fn invalidate_for_weight_change(&mut self) {
3369        self.clear_sequence_state();
3370        self.embryo_graph = None;
3371    }
3372
3373    /// Prompt positions the cache already holds when `input_ids`
3374    /// strictly EXTENDS the consumed prefix (0 otherwise). Bounded-native
3375    /// models answer from the fixed-size `kv_prefix` record; everything
3376    /// else from the legacy `kv_history` vector (which the network split
3377    /// also reads and writes).
3378    fn cached_prefix_len(&self, input_ids: &[u32]) -> usize {
3379        let (n, on_device) = if self.bounded_native() {
3380            (self.kv_prefix.extension(input_ids), self.kv_prefix.on_device())
3381        } else {
3382            let h = &self.kv_history;
3383            if !h.is_empty() && h.len() < input_ids.len() && input_ids[..h.len()] == h[..] {
3384                (h.len(), self.kv_history_device)
3385            } else {
3386                (0, false)
3387            }
3388        };
3389        // The owner tag: the state the key describes must live where this
3390        // turn will continue it (R4) — else a fresh sequence.
3391        if n > 0 && !self.prefix_owner_matches(n, on_device) {
3392            tracing::warn!(
3393                "kv-reuse refused: the cached prefix ({n} positions) was built on the {} path, \
3394                 the device now holds {:?} — re-prefilling from zero",
3395                if on_device { "resident device" } else { "host" },
3396                self.device_sequence_position()
3397            );
3398            return 0;
3399        }
3400        n
3401    }
3402
3403    /// Public view of the prefix-reuse decision for `input_ids` (positions
3404    /// the next `generate*` would take from the cache; 0 = fresh sequence).
3405    pub fn reusable_prefix_len(&self, input_ids: &[u32]) -> usize {
3406        self.cached_prefix_len(input_ids)
3407    }
3408
3409    /// Record the forwarded prefix as the next turn's reuse key. A
3410    /// bounded-native model extends the rolling record (`reused` = the
3411    /// positions this turn found cached); others keep the literal vector.
3412    /// Either way the key carries its OWNER: the resident device graph
3413    /// (it holds an image of this sequence) or the host.
3414    fn record_consumed_prefix(&mut self, consumed: &[u32], reused: usize) {
3415        let on_device = self.device_sequence_position().is_some();
3416        if self.bounded_native() {
3417            let keep = reused > 0 && reused == self.kv_prefix.len() && reused <= consumed.len();
3418            let prev_device = self.kv_prefix.on_device();
3419            self.kv_history.clear();
3420            self.kv_history_device = false;
3421            if keep && prev_device == on_device {
3422                self.kv_prefix.extend(&consumed[reused..]);
3423            } else {
3424                self.kv_prefix.set(consumed);
3425            }
3426            self.kv_prefix.set_on_device(on_device);
3427        } else {
3428            self.kv_history = consumed.to_vec();
3429            self.kv_history_device = on_device;
3430        }
3431    }
3432
3433    /// True when at least one layer runs the O(1) kernel.
3434    pub fn o1_active(&self) -> bool {
3435        self.o1_cfg.is_some() && self.o1_flags.iter().any(|&f| f)
3436    }
3437
3438    /// Whether generation's prompt ingest is routed through the whole-token
3439    /// graph.  The bench uses this to label the measured generation prefill
3440    /// honestly; keep the predicate in Pipeline so CLI labels cannot drift
3441    /// from the production route.
3442    /// Positions per batched-graph submit for the prompt: `CMF_BATCH_K`
3443    /// when set (0 = one position at a time through the token graph),
3444    /// otherwise 32 on a discrete card whose prompt takes the graph route.
3445    /// The batched graph read a 2048-token prompt at 53 tok/s against 28.5
3446    /// one position at a time on an RTX PRO 4000 (Qwen3.8-27B q4tp: TTFT
3447    /// 39 s against 72), and its states are the speculative verify's,
3448    /// measured identical to the plain path. macOS keeps its own arm.
3449    pub fn generation_batch_k(&self) -> usize {
3450        if let Some(k) = std::env::var("CMF_BATCH_K")
3451            .ok()
3452            .and_then(|v| v.parse::<usize>().ok())
3453        {
3454            return k;
3455        }
3456        #[cfg(not(target_os = "macos"))]
3457        if self.graph_prefill_preferred() && !self.o1_active() {
3458            return 32;
3459        }
3460        0
3461    }
3462
3463    pub fn generation_graph_prefill(&self) -> bool {
3464        let graph = self.graph_prefill_preferred();
3465        // On wgpu, an active MTP head now consumes the trunk's graph batches
3466        // and warms its own block from those returned rows.  The selected
3467        // generation measurement is therefore the batched path, even though
3468        // the underlying GDN model still satisfies the graph-prefill
3469        // predicate.  Keep the CLI label tied to the actual route.  Native
3470        // Metal has a separate prefill-batch arm and retains its historical
3471        // label here.
3472        // A batched prompt (`generation_batch_k` > 0) is the batched graph
3473        // for every model on the graph route, not only those with an MTP
3474        // head — the label follows the route.
3475        #[cfg(not(target_os = "macos"))]
3476        if graph
3477            && self.generation_batch_k() > 0
3478            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
3479        {
3480            return false;
3481        }
3482        graph
3483    }
3484
3485    /// Device-side O(1) mirrors currently uploaded for this pipeline's
3486    /// sequence.  The count/bytes are zero before seal or after a fresh
3487    /// reset; callers use this to distinguish logical host state from the
3488    /// GPU allocation that actually serves decode.
3489    pub fn o1_device_stats(&self) -> (usize, u64) {
3490        crate::gpu::o1_device_stats(self.graph_kv_id)
3491    }
3492
3493    /// Arm query collection on the o1 layers (fresh prompt pass).
3494    /// Reset the o1 layers to Collecting for a fresh sequence. Pub for the
3495    /// network split: each side runs the o1 lifecycle over ITS OWN layers
3496    /// (begin before prefill, seal at the prefill barrier).
3497    pub fn o1_begin(&mut self) {
3498        self.o1_begin_with_prefix(None);
3499    }
3500
3501    /// Arm collection and optionally request a positive calibration prefix.
3502    /// The effective barrier is always at least the skeleton-safe floor, so
3503    /// a short requested prefix cannot create an exact-only runtime state.
3504    pub fn o1_begin_with_prefix(&mut self, requested_prefix: Option<usize>) {
3505        if let Some(c) = &self.o1_cfg {
3506            let (m, w, sink, rect) = (c.m, c.w, c.sink, c.rect);
3507            let boundary = requested_prefix.map(|p| {
3508                p.max(
3509                    crate::nystrom::o1_deferred_boundary(w, sink)
3510                        .expect("o1 config boundary validated in set_o1"),
3511                )
3512            });
3513            for (li, &f) in self.o1_flags.iter().enumerate() {
3514                if f {
3515                    self.kv_cache.layers[li].o1_begin_with_boundary(m, w, sink, rect, boundary);
3516                }
3517            }
3518        }
3519    }
3520
3521    /// Effective deferred boundary for a positive prefix request.
3522    fn o1_effective_boundary(&self, requested_prefix: usize) -> Option<usize> {
3523        self.o1_cfg.as_ref().and_then(|c| {
3524            crate::nystrom::o1_deferred_boundary(c.w, c.sink)
3525                .map(|floor| requested_prefix.max(floor))
3526        })
3527    }
3528
3529    fn o1_note_transition(&mut self) {
3530        // Drain every layer's one-shot bit before publishing one pipeline
3531        // epoch. `any()` would short-circuit on the first layer and leak the
3532        // remaining bits into later forwards, causing one epoch per layer.
3533        let mut transitioned = false;
3534        for (li, &flagged) in self.o1_flags.iter().enumerate() {
3535            if flagged {
3536                transitioned |= self.kv_cache.layers[li].take_o1_transition();
3537            }
3538        }
3539        if transitioned {
3540            self.o1_epoch = self.o1_epoch.wrapping_add(1);
3541        }
3542    }
3543
3544    fn o1_pending(&self) -> bool {
3545        self.o1_flags.iter().enumerate().any(|(li, &f)| {
3546            f && self.kv_cache.layers[li].seq_len > 0
3547                && self.kv_cache.layers[li].o1_pending_boundary().is_some()
3548        })
3549    }
3550
3551    fn o1_fail(&mut self, err: String) {
3552        tracing::error!("o1 deferred seal failed; terminating sequence: {err}");
3553        self.clear_sequence_state();
3554        self.graph_failed
3555            .store(true, std::sync::atomic::Ordering::Relaxed);
3556        self.cancel
3557            .store(true, std::sync::atomic::Ordering::Relaxed);
3558    }
3559
3560    /// Seal participating layers while retaining the exact state when the
3561    /// prompt is below the deferred boundary. A split worker may have
3562    /// collecting layers outside its owned span; zero-depth layers remain
3563    /// armed and are intentionally skipped until their peer runs them.
3564    pub fn o1_seal_checked(&mut self) -> Result<bool, String> {
3565        if self.o1_cfg.is_none() {
3566            return Ok(false);
3567        }
3568        let mut participating = false;
3569        for li in 0..self.num_layers {
3570            if !self.o1_flags.get(li).copied().unwrap_or(false) {
3571                continue;
3572            }
3573            if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3574                return Err(err);
3575            }
3576            if self.kv_cache.layers[li].seq_len == 0 {
3577                continue;
3578            }
3579            participating = true;
3580            let num_heads = self.layer_num_heads(li);
3581            self.kv_cache.layers[li].o1_seal_checked(num_heads)?;
3582        }
3583        self.o1_note_transition();
3584        for li in 0..self.num_layers {
3585            if self.o1_flags.get(li).copied().unwrap_or(false) {
3586                if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3587                    return Err(err);
3588                }
3589            }
3590        }
3591        Ok(participating
3592            && (0..self.num_layers).all(|li| {
3593                !self.o1_flags.get(li).copied().unwrap_or(false)
3594                    || self.kv_cache.layers[li].seq_len == 0
3595                    || self.kv_cache.layers[li].o1_sealed()
3596            }))
3597    }
3598
3599    /// Complete a deferred boundary after a full position/span forward.
3600    /// This is the pipeline owner for epoch publication and failure cleanup.
3601    fn o1_progress(&mut self) {
3602        if !self.o1_active() {
3603            return;
3604        }
3605        for li in 0..self.num_layers {
3606            if self.o1_flags.get(li).copied().unwrap_or(false) {
3607                if let Some(err) = self.kv_cache.layers[li].take_o1_error() {
3608                    self.o1_fail(err);
3609                    return;
3610                }
3611            }
3612        }
3613        // A qwen_attention row can seal in the middle of a complete layer
3614        // walk. Consume its transition even though the pending boundary has
3615        // already disappeared from the cache.
3616        self.o1_note_transition();
3617        if !self.o1_pending() {
3618            return;
3619        }
3620        if let Err(err) = self.o1_seal_checked() {
3621            self.o1_fail(err);
3622        }
3623    }
3624
3625    /// Turn a deferred O(1) failure raised by a hidden-only forward into the
3626    /// Result error its public batch/span caller must return. The failure
3627    /// path already cleared host/device sequence state; consume only the
3628    /// side-channel marker here and leave the pipeline reusable.
3629    fn check_o1_progress_failure(&mut self, phase: &str) -> Result<(), String> {
3630        if self
3631            .graph_failed
3632            .swap(false, std::sync::atomic::Ordering::Relaxed)
3633        {
3634            self.cancel
3635                .store(false, std::sync::atomic::Ordering::Relaxed);
3636            self.clear_sequence_state();
3637            return Err(format!("{phase}: deferred O(1) transition failed"));
3638        }
3639        Ok(())
3640    }
3641
3642    /// Freeze landmarks + skeleton state after the prompt pass and drop
3643    /// the o1 layers' full KV; decode then runs `step()` per token.
3644    /// Pub for the network split (see `o1_begin`).
3645    pub fn o1_seal(&mut self) {
3646        if let Err(err) = self.o1_seal_checked() {
3647            self.o1_fail(err);
3648        }
3649    }
3650
3651    /// Enable/disable the structured per-token telemetry trace (B4).
3652    pub fn set_trace(&mut self, on: bool) {
3653        self.trace = on;
3654    }
3655
3656    /// Replace all request-scoped sampler options and reset the random stream.
3657    /// This is required for deterministic `seed` semantics in pooled servers.
3658    pub fn set_sampler_config(&mut self, config: SamplerConfig) {
3659        self.rng = match config.seed {
3660            Some(seed) => SplitMix64::new(seed),
3661            None => SplitMix64::from_entropy(),
3662        };
3663        self.sampler_config = config;
3664    }
3665
3666    /// Toggle the per-token confidence reduction (a full-vocab
3667    /// softmax each token). `bench --core` turns it off so the timed
3668    /// loop matches llama-bench's core contract; the result's
3669    /// `confidence` vec is empty while off.
3670    pub fn set_confidence(&mut self, on: bool) {
3671        self.confidence_on = on;
3672    }
3673
3674    /// Set the confidence-calibration temperature (B1). Values ≤0 are
3675    /// clamped to raw (1.0).
3676    pub fn set_calib_temp(&mut self, t: f32) {
3677        self.calib_temp = if t > 1e-3 { t } else { 1.0 };
3678    }
3679
3680    /// The active calibration temperature (1.0 = raw probability).
3681    pub fn calib_temp(&self) -> f32 {
3682        self.calib_temp
3683    }
3684
3685    /// Partial rotary (Qwen3.5): rotate only the first `rotary_dim` dims;
3686    /// the frequency table is rebuilt over the rotary dims.
3687    pub fn set_rotary(&mut self, rotary_dim: usize, base: f32) {
3688        self.rotary_dim = rotary_dim.min(self.head_dim);
3689        self.inv_freq = std::sync::Arc::new(attention::rope_inv_freq(self.rotary_dim, base));
3690        // The packed resident graph owns its own inverse-frequency plane;
3691        // changing RoPE after it was built must not leave a stale device
3692        // model behind the exact host configuration.
3693        self.embryo_graph = None;
3694    }
3695
3696    fn attn_cfg(&self, position: usize) -> QwenAttnCfg<'_> {
3697        QwenAttnCfg {
3698            num_heads: self.num_heads,
3699            num_kv_heads: self.num_kv_heads,
3700            head_dim: self.head_dim,
3701            hidden_size: self.hidden_size,
3702            position,
3703            inv_freq: &self.inv_freq,
3704            rotary_dim: self.rotary_dim,
3705            scale: self.attn_scale,
3706            softcap: self.attn_softcap,
3707            window: None,
3708            v_norm: false,
3709            qk_norm_after_rope: self.qk_norm_after_rope,
3710            gate_sigmoid: self.proj_gate_sigmoid,
3711            q_norm: None,
3712            k_norm: None,
3713            output_gate: false,
3714            softplus_gate: None,
3715            rope_scale: self.rope_scale,
3716            bias: None,
3717            rms_eps: self.rms_eps,
3718            norm_style: self.norm_style,
3719            pool: self.pool.as_deref(),
3720            v_head_dim: self.v_head_dim.unwrap_or(self.head_dim),
3721        }
3722    }
3723
3724    /// Generate text from a plain-text prompt. Streams tokens via `on_token`.
3725    pub fn generate(
3726        &mut self,
3727        prompt: &str,
3728        max_tokens: usize,
3729        task_mask: Option<&TaskMask>,
3730        on_token: Option<TokenCallback>,
3731    ) -> Result<GenerateResult, String> {
3732        let input_ids = self.tokenizer.with_bos(self.tokenizer.encode(prompt));
3733        self.generate_from_ids(&input_ids, max_tokens, task_mask, on_token)
3734    }
3735
3736    /// Generate from a V4.1 multimodal prompt prepared by the vision module.
3737    /// Vision rows are encoded once and fed through the same bounded token walk as text.
3738    pub fn generate_from_vl(
3739        &mut self,
3740        input: &crate::dsv41_vision::PreparedVlInputs,
3741        max_tokens: usize,
3742        task_mask: Option<&TaskMask>,
3743        on_token: Option<TokenCallback>,
3744    ) -> Result<GenerateResult, String> {
3745        let Some(dsv41) = &self.dsv41 else {
3746            return Err("V4.1 multimodal input requires a DeepSeek-V4.1 pipeline".into());
3747        };
3748        if input.token_ids.is_empty() {
3749            return Err("empty V4.1 multimodal prompt".into());
3750        }
3751        if input.token_types.len() != input.token_ids.len() {
3752            return Err(format!(
3753                "V4.1 token type count {} != token count {}",
3754                input.token_types.len(),
3755                input.token_ids.len()
3756            ));
3757        }
3758        let dim = dsv41.2.dim;
3759        let mut embeddings = vec![None; input.token_ids.len()];
3760        let mut participates = vec![true; input.token_ids.len()];
3761        if !input.images.is_empty() {
3762            let vision = self
3763                .dsv41_vision
3764                .as_ref()
3765                .ok_or_else(|| "V4.1 image prompt has no loaded vision tower".to_string())?;
3766            for image in &input.images {
3767                let end = image.start.saturating_add(image.types.len());
3768                if end > input.token_ids.len() {
3769                    return Err(format!(
3770                        "V4.1 image span {}..{} exceeds prompt length {}",
3771                        image.start,
3772                        end,
3773                        input.token_ids.len()
3774                    ));
3775                }
3776                let mut span = vec![0.0f32; image.types.len() * dim];
3777                vision.fill_image_span(image, &mut span, self.pool.as_deref())?;
3778                for (offset, &kind) in image.types.iter().enumerate() {
3779                    let pos = image.start + offset;
3780                    if input.token_types[pos] != kind {
3781                        return Err(format!(
3782                            "V4.1 image type mismatch at position {pos}: {} != {kind}",
3783                            input.token_types[pos]
3784                        ));
3785                    }
3786                    embeddings[pos] = Some(span[offset * dim..(offset + 1) * dim].to_vec());
3787                    participates[pos] = false;
3788                }
3789            }
3790        }
3791        for (pos, &kind) in input.token_types.iter().enumerate() {
3792            if kind == crate::dsv41_vision::TEXT && embeddings[pos].is_some() {
3793                return Err(format!("V4.1 text position {pos} has an image embedding"));
3794            }
3795            if kind != crate::dsv41_vision::TEXT && embeddings[pos].is_none() {
3796                return Err(format!("V4.1 image position {pos} has no image embedding"));
3797            }
3798        }
3799        self.dsv41_prefill = Some((embeddings, participates));
3800        let result = self.generate_from_ids(&input.token_ids, max_tokens, task_mask, on_token);
3801        self.dsv41_prefill = None;
3802        result
3803    }
3804
3805    /// `None` when the mask forbids nothing (see `TaskMask::fully_open`).
3806    fn drop_open_mask<'m>(&self, m: Option<&'m TaskMask>) -> Option<&'m TaskMask> {
3807        m.filter(|m| !m.fully_open(self.intermediate_size, self.num_heads))
3808    }
3809
3810    /// Generate from prepared token ids (e.g. a chat template).
3811    ///
3812    /// With an MTP head, greedy generation without a task mask takes the
3813    /// speculative path: the MTP module drafts the token after next and
3814    /// the main model verifies both in one fused two-position forward
3815    /// (weights streamed once). The output is EXACTLY the vanilla greedy
3816    /// sequence — a rejected draft is rolled back — MTP only buys speed.
3817    pub fn generate_from_ids(
3818        &mut self,
3819        input_ids: &[u32],
3820        max_tokens: usize,
3821        task_mask: Option<&TaskMask>,
3822        on_token: Option<TokenCallback>,
3823    ) -> Result<GenerateResult, String> {
3824        self.generate_with_prompt_rows(input_ids, None, max_tokens, task_mask, on_token)
3825    }
3826
3827    /// Generate from complete prompt embeddings [token_count, hidden_size].
3828    /// Text rows can be obtained with `embed_id`; media rows replace only
3829    /// their expanded placeholder positions. Rows are already scaled and
3830    /// enter `PrefillIn::Hidden`, so a device graph must not re-embed them.
3831    /// Token-only KV reuse is disabled both into and out of this request.
3832    pub fn generate_from_embeds(
3833        &mut self,
3834        input_ids: &[u32],
3835        prompt_rows: &[f32],
3836        max_tokens: usize,
3837        task_mask: Option<&TaskMask>,
3838        on_token: Option<TokenCallback>,
3839    ) -> Result<GenerateResult, String> {
3840        if input_ids.is_empty()
3841            || input_ids.len().checked_mul(self.hidden_size) != Some(prompt_rows.len())
3842        {
3843            return Err("embedded prompt dimensions must be [tokens, hidden_size]".into());
3844        }
3845        if prompt_rows.iter().any(|x| !x.is_finite()) {
3846            return Err("embedded prompt contains non-finite values".into());
3847        }
3848        if !self.can_prefill_batched() || self.dyn_router.is_some()
3849            || self.o1_active() || self.mtp.is_some() || self.gpu_plan.is_some()
3850        {
3851            return Err("embedded prompts require the ordinary transformer path without O(1), dynamic routing, GPU splitting or a generic MTP head".into());
3852        }
3853        self.generate_with_prompt_rows(input_ids, Some(prompt_rows), max_tokens, task_mask, on_token)
3854    }
3855
3856    fn generate_with_prompt_rows(
3857        &mut self,
3858        input_ids: &[u32],
3859        prompt_rows: Option<&[f32]>,
3860        max_tokens: usize,
3861        task_mask: Option<&TaskMask>,
3862        mut on_token: Option<TokenCallback>,
3863    ) -> Result<GenerateResult, String> {
3864        #[cfg(target_os = "macos")]
3865        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
3866        if std::env::var("CMF_TRACE_H").is_ok() {
3867            eprintln!("input_ids: {input_ids:?}");
3868        }
3869        if input_ids.is_empty() {
3870            return Err("empty prompt: nothing to generate from".to_string());
3871        }
3872        // A prior graph failure is terminal for that sequence but must not
3873        // poison the next independent request.  Keep this flag separate from
3874        // the externally-owned cooperative cancel bit.
3875        self.graph_failed
3876            .store(false, std::sync::atomic::Ordering::Relaxed);
3877        // A mask that forbids nothing still costs every fused path and
3878        // whole-token graph, all of which are gated on `is_none()`. A
3879        // narrowed file whose one segment is always on carries exactly
3880        // such a mask — drop it here rather than pay 5x for a no-op.
3881        let task_mask = self.drop_open_mask(task_mask);
3882
3883        // Cross-turn KV reuse: a chat app resends the whole history
3884        // every turn; when the new ids strictly EXTEND what the cache
3885        // already holds, prefill only the tail — turn latency stays
3886        // proportional to the new text instead of the whole session.
3887        // Extension-only (no rollback), so it is exact for every layer
3888        // kind including recurrent state; MTP/o1/task-mask runs keep
3889        // the fresh-sequence path. CMF_KV_REUSE=0 disables.
3890        let mut reuse_from = {
3891            let on = !std::env::var("CMF_KV_REUSE").is_ok_and(|v| v == "0");
3892            if on
3893                && prompt_rows.is_none()
3894                && task_mask.is_none()
3895                && self.mtp.is_none()
3896                && !(self.mimo_mtp.is_some() && self.speculative)
3897                && self.o1_cfg.is_none()
3898                && self.dsv41.is_none()
3899            {
3900                self.cached_prefix_len(input_ids)
3901            } else {
3902                0
3903            }
3904        };
3905        // The device may own rows the host tail prefill needs (wgpu decode
3906        // writes only its mirror): hand them to the host, or start fresh.
3907        if reuse_from > 0 && !self.prepare_kv_reuse(reuse_from) {
3908            reuse_from = 0;
3909        }
3910        self.last_prefill_tokens = input_ids.len() - reuse_from;
3911        let bounded_native = self.bounded_native();
3912        if reuse_from == 0 {
3913            // Fresh sequence — the cache holds absolute positions.
3914            self.clear_sequence_state();
3915        } else if std::env::var("CMF_PREFILL_PROF").is_ok() {
3916            eprintln!(
3917                "kv-reuse: {} of {} prompt positions already cached",
3918                reuse_from,
3919                input_ids.len()
3920            );
3921        }
3922        crate::gpu::graph_race_begin_generation();
3923        // Optional bounded calibration prefix. Keep the requested value
3924        // even when it is longer than the prompt; the collecting layer will
3925        // defer at the effective boundary and remain exact for short input.
3926        let o1_prefill = if self.o1_active() && task_mask.is_none() {
3927            std::env::var("CMF_O1_PREFILL")
3928                .ok()
3929                .and_then(|v| v.parse::<usize>().ok())
3930                .filter(|&p| p > 0)
3931        } else {
3932            None
3933        };
3934        if task_mask.is_none() {
3935            self.o1_begin_with_prefix(o1_prefill);
3936        }
3937
3938        // Speculative decode is off under o1: a rejected draft can't be
3939        // rolled back out of the far accumulators / ring window (the
3940        // Nyström insertion is irreversible by design).
3941        // The wgpu token graph owns a device K/V mirror that speculative
3942        // rollback would desync — the two are mutually exclusive.
3943        let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
3944        // Graph speculative decode (`CMF_GRAPH_SPEC=1`): the MTP head
3945        // drafts, ONE batched graph submit verifies the whole chain.
3946        //
3947        // It now PAYS on Qwen3.6-27B / RTX 5090 — 51.1 tok/s against a
3948        // plain 49.4 at k=3, medians of three, 89% of drafts accepted,
3949        // and the greedy continuation is byte-identical to the plain
3950        // path. That took the batch matvec sharing its nibble unpack
3951        // across the batch (`CMF_MV_BK=2`); before it, the same round
3952        // measured 43.6, an 11% LOSS, which is what the earlier note
3953        // here described.
3954        //
3955        // Still opt-in. One model's win is not a default: the verify
3956        // rides `gdn_spec_restore` and a batched frame whose numerics
3957        // are the batch kernels', and that has to be shown on more than
3958        // one architecture before every greedy decode takes it.
3959        // Greedy (with or without penalties) verifies by argmax equality.
3960        // Sampling (temperature > 0) can go through speculative SAMPLING —
3961        // draft from the MTP head's own post-chain distribution, accept
3962        // with min(1, p/q), correct from max(0, p − q); the emitted stream
3963        // is distributed exactly as the plain sampler's — but it is
3964        // OPT-IN (`CMF_GRAPH_SPEC_SAMPLE=1`): measured on Qwen3.8-27B /
3965        // RTX 5090 at the instruct row (0.7 / 0.80 / 20 / presence 1.5)
3966        // it decoded 19-22 tok/s against a plain 40 — nine post-chain
3967        // distributions a round plus a lower acceptance than greedy's,
3968        // against a verify that costs 2.7 single tokens. The greedy arms
3969        // pay +10%; the sampling arm needs a cheaper verify first.
3970        // Native Metal HAS that verify: its eight-row tile is flat in b,
3971        // so a round costs ~1.9 plain tokens and the sampling arm pays at
3972        // 2.3 accepted per round — measured on Qwen3.8-27B q4tp / M4 at
3973        // the CLI defaults (0.7 / rep 1.1 / top-k 40, seed 42), a code
3974        // prompt: 9.0 tok/s against a plain 5.4 in the same window, and
3975        // the per-round watchdog turns it off where prose loses. So on
3976        // Metal the sampling arm is ON (`CMF_GRAPH_SPEC_SAMPLE=0` opts out)
3977        // — but only for a config the SPARSE chain serves (a top-k within
3978        // `sparse_ok`): without it a round builds nine 248k-float
3979        // distributions on the host, which is the 5090's measured loss and
3980        // not a cost the round-token proxy below can see. A top-k-less
3981        // sampling config keeps the plain path unless asked for by name.
3982        #[cfg(target_os = "macos")]
3983        let metal_graph = crate::gpu::q1_force()
3984            && crate::gpu::enabled_here()
3985            && std::env::var("CMF_GPU_BLOCK")
3986                .map(|v| v != "0")
3987                .unwrap_or(true);
3988        #[cfg(not(target_os = "macos"))]
3989        let metal_graph = false;
3990        let spec_sample_env = std::env::var("CMF_GRAPH_SPEC_SAMPLE").ok();
3991        // A round whose cost is the MEASURED one: greedy (argmax rows), or
3992        // sampling through the sparse chain. Anything else pays the dense
3993        // chain's host time, which no proxy can price.
3994        let spec_cheap_round = self.sampler_config.temperature < 1e-6
3995            || sampler::sparse_ok(&self.sampler_config);
3996        let spec_sampling_ok = self.sampler_config.temperature < 1e-6
3997            || match spec_sample_env.as_deref() {
3998                Some("1") => true,
3999                Some(_) => false,
4000                None => metal_graph && spec_cheap_round,
4001            };
4002        // ON by default for greedy on the wgpu graph: with the draft on
4003        // the graph and the verify bit-exact, it measured 58.7 tok/s
4004        // against a plain 48.1 on Qwen3.8-27B q4tp / RTX 5090 (k=4) and
4005        // 51.1 against 49.4 on Qwen3.6-27B, and a round that stops
4006        // paying turns itself off below (acceptance watchdog).
4007        // `CMF_GRAPH_SPEC=0` disables; `=1` was the old opt-in spelling.
4008        // …but only where the batched verify has its register-blocked
4009        // kernel: q4tp dense FFNs (graph kind 6). q4t and q8_2f verify
4010        // through tile GEMMs today and measured a LOSS (q8_2f 22 against
4011        // 29 tok/s), the 2-bit plane the same; those stay opt-in
4012        // (`CMF_GRAPH_SPEC=1`).
4013        // …at least in nine dense FFNs of ten: a healed file carries its
4014        // last two layers at q8_2f, and two tile-GEMM verifies among 64 do
4015        // not change the arithmetic (measured: the healed q4tp file
4016        // decodes at the plain file's rate and would otherwise sit out).
4017        let (mut dense_n, mut dense_q4tp) = (0usize, 0usize);
4018        for lw in &self.weights.layers {
4019            if let FfnKind::Dense(d) = &lw.ffn {
4020                dense_n += 1;
4021                if matches!(d.gate_proj.graph_weight(), Some((_, _, 6, _)))
4022                    && matches!(d.up_proj.graph_weight(), Some((_, _, 6, _)))
4023                    && matches!(d.down_proj.graph_weight(), Some((_, _, 6, _)))
4024                {
4025                    dense_q4tp += 1;
4026                }
4027            }
4028        }
4029        let spec_default_ok = dense_n == 0 || dense_q4tp * 10 >= dense_n * 9;
4030        // Penalties break the draft head's agreement with the trunk (a
4031        // 1.1 repetition penalty measured 2 of 16 accepted): not by
4032        // default there either — off Metal that rule is untouched, and
4033        // suppressed ids keep counting as a penalty there, because no
4034        // measurement on a discrete card says otherwise.
4035        //
4036        // On native Metal the penalized arms DO pay: the draft applies
4037        // the same penalty and the verify scores the penalized rows
4038        // exactly (`greedy_pen`, the plain loop's arithmetic), so the
4039        // text is the plain path's and only the round's shape changes.
4040        // Measured on this M4 — see the report for the interleaved run.
4041        let penalized = !metal_graph
4042            && (self.sampler_config.repetition_penalty != 1.0
4043                || self.sampler_config.presence_penalty != 0.0
4044                || !self.sampler_config.suppress_tokens.is_empty());
4045        // …and not on wgpu-over-Metal: the batched verify graph there
4046        // returned 0 accepted drafts and garbage text on a GDN hybrid
4047        // (16.08, Qwen3.5-0.8B) while Vulkan is bit-exact; the Mac's
4048        // default backend is native Metal without a batch graph anyway.
4049        #[cfg(feature = "gpu")]
4050        let metal_wgpu = graph_on && crate::gpu_wgpu::wgpu_backend_is_metal();
4051        #[cfg(not(feature = "gpu"))]
4052        let metal_wgpu = false;
4053        let spec_env = std::env::var("CMF_GRAPH_SPEC").ok();
4054        let spec_wanted = match spec_env.as_deref() {
4055            Some("0") => false,
4056            Some(_) => {
4057                if metal_wgpu {
4058                    tracing::warn!(
4059                        "CMF_GRAPH_SPEC forced on wgpu/Metal: the batched verify graph is not \
4060                         verified on this backend (garbage measured on Qwen3.5-0.8B)"
4061                    );
4062                }
4063                true
4064            }
4065            None => spec_default_ok && !penalized && !metal_wgpu,
4066        };
4067        // Native Metal: the b-row verify graph (`try_batch_graph_metal`)
4068        // stands where the wgpu batch graph stands on discrete cards
4069        // (`metal_graph`, above).
4070        let graph_spec = self.speculative
4071            && (graph_on || metal_graph)
4072            && self.mtp.is_some()
4073            && task_mask.is_none()
4074            && !self.o1_active()
4075            && spec_sampling_ok
4076            && spec_wanted;
4077        // Native Metal: say the route ONCE (RUST_LOG=info), so a user can
4078        // confirm the fast path without setting a single flag — every
4079        // knob below defaults to the measured-best value on the M4.
4080        #[cfg(target_os = "macos")]
4081        if metal_graph {
4082            static SAID: std::sync::Once = std::sync::Once::new();
4083            SAID.call_once(|| {
4084                let spec = if graph_spec {
4085                    let k = std::env::var("CMF_GRAPH_SPEC_K")
4086                        .ok()
4087                        .and_then(|v| v.parse::<usize>().ok())
4088                        .filter(|&v| (1..=8).contains(&v))
4089                        .unwrap_or(7);
4090                    let arm = if self.sampler_config.temperature < 1e-6 {
4091                        "greedy"
4092                    } else {
4093                        "sampling"
4094                    };
4095                    format!(
4096                        "spec k={k} {arm} (batched verify, draft shortlist {}, trial: proxy)",
4097                        Self::draft_vocab_rows(usize::MAX)
4098                    )
4099                } else if !self.speculative {
4100                    "spec off (CMF_MTP=0)".to_string()
4101                } else if self.mtp.is_none() {
4102                    "spec off (no MTP head)".to_string()
4103                } else if !spec_sampling_ok {
4104                    if spec_cheap_round {
4105                        "spec off (CMF_GRAPH_SPEC_SAMPLE=0)".to_string()
4106                    } else {
4107                        "spec off (sampling without a top-k: the dense chain \
4108                         costs more than it saves)"
4109                            .to_string()
4110                    }
4111                } else if !spec_wanted {
4112                    "spec off (CMF_GRAPH_SPEC=0 or non-q4tp FFNs)".to_string()
4113                } else if task_mask.is_some() {
4114                    "spec off (task mask)".to_string()
4115                } else {
4116                    "spec off (O(1) attention)".to_string()
4117                };
4118                let on = |var: &str| {
4119                    if std::env::var(var).as_deref() == Ok("0") {
4120                        "off"
4121                    } else {
4122                        "on"
4123                    }
4124                };
4125                tracing::info!(
4126                    "metal native: {spec}, state4 {}, async replay {}, prefill graph {}, \
4127                     MTP graph {}, attend {}, probe {}",
4128                    if crate::gpu_metal::state4_on() { "on" } else { "off" },
4129                    if crate::gpu_metal::async_replay_on() { "on" } else { "off" },
4130                    on("CMF_METAL_PREFILL"),
4131                    on("CMF_MTP_GRAPH"),
4132                    std::env::var("CMF_GPU_ATTEND").unwrap_or_else(|_| "auto".into()),
4133                    if crate::gpu::probe_enabled() { "bypassed (q1 force)" } else { "off" },
4134                );
4135            });
4136        }
4137        // GDN hybrids sit the fused-pair speculation out by default: the
4138        // recurrence is sequential, so the pair lane cannot parallelize
4139        // (the bench's own Pair line reads fused 1.28x TWO singles on the
4140        // 35B) and the draft's full-vocab head rides on top — measured 2x
4141        // SLOWER end to end (16.1 vs 32.4 tok/s on the 48-core stand).
4142        // CMF_MTP=1 forces it back for study.
4143        let pair_pays = self.gdn_cfg.is_none() || std::env::var("CMF_MTP").as_deref() == Ok("1");
4144        let spec_active = self.speculative
4145            && self.mtp.is_some()
4146            && task_mask.is_none()
4147            && !self.o1_active()
4148            && ((!graph_on && pair_pays && self.sampler_config.temperature < 1e-6) || graph_spec);
4149        // The MTP module is detached during generation so its mutable
4150        // state does not fight the borrow on `self`.
4151        let mut mtp = if spec_active { self.mtp.take() } else { None };
4152        if std::env::var("CMF_MTP_CHAIN_PROBE").is_ok() {
4153            eprintln!(
4154                "mtp-probe gate: spec_active={spec_active} mtp={} speculative={} graph_on={graph_on} temp_ok={}",
4155                mtp.is_some(),
4156                self.speculative,
4157                self.sampler_config.temperature < 1e-6,
4158            );
4159        }
4160        if let Some(m) = &mut mtp {
4161            m.kv.clear();
4162            // The MTP block's own device mirror starts over with its cache.
4163            crate::gpu::graph_kv_reset(self.mtp_kv_id());
4164            self.mtp_graph_mode = None;
4165        }
4166        // MiMo-V2's draft stack: greedy rounds (draft K with the chained
4167        // MTP layers, verify K+1 rows in one batched forward). Sampling
4168        // decodes plain; `CMF_MTP=0` / `CMF_MIMO_MTP=0` turn it off.
4169        let mimo_spec = self.speculative
4170            && self.mimo_mtp.is_some()
4171            && task_mask.is_none()
4172            && !self.o1_active()
4173            && self.dyn_router.is_none()
4174            && self.sampler_config.temperature < 1e-6
4175            && std::env::var("CMF_MIMO_MTP").as_deref() != Ok("0");
4176        if let Some(st) = self.mimo_mtp.as_mut() {
4177            st.reset();
4178            if mimo_spec && std::env::var_os("CMF_MIMO_MTP_PROBE").is_some() {
4179                Self::mimo_mtp_hist_cap(st, input_ids.len());
4180            }
4181        }
4182        // Dynamic router detached during decode (same borrow trick as MTP).
4183        // Speculative decode and dynamic routing are mutually exclusive
4184        // for now — the fused-pair path doesn't carry per-token φ.
4185        let mut router = if mtp.is_none() {
4186            self.dyn_router.take()
4187        } else {
4188            None
4189        };
4190        let mut reuse_from = reuse_from;
4191        if let Some(r) = &mut router {
4192            r.reset(); // active=backbone, matching a fresh overlay
4193            self.dyn_phi_seen = 0; // fresh φ EMA per generation
4194            if self.dyn_active.is_some() {
4195                // A real switch back to the backbone invalidates the
4196                // cache the reuse key was computed against.
4197                let _ = self.set_active_skill(None);
4198                reuse_from = 0;
4199                self.last_prefill_tokens = input_ids.len();
4200            }
4201        }
4202
4203        let mut all_ids = input_ids.to_vec();
4204        let mut generated = 0usize;
4205        let mut finish_reason = "max_tokens".to_string();
4206        let mut drafted = 0usize;
4207        let mut accepted = 0usize;
4208        // DeepSeek-V4's draft quality is strongly content-dependent.  Two
4209        // consecutive paid rounds with no extra token put it on a bounded
4210        // cooldown; predictable text keeps batching, ordinary prose falls
4211        // back to the exact walk instead of paying a slow draft forever.
4212        // Local to one generation so one difficult request cannot poison the
4213        // next one, and deliberately automatic — this is not a user knob.
4214        let mut dsv4_spec_bad = 0usize;
4215        let mut dsv4_spec_retry_at = 0usize;
4216        let mut confidence: Vec<f32> = Vec::new();
4217        let trace_on = self.trace;
4218        let calib_temp = self.calib_temp;
4219        let mut traces: Vec<TokenTrace> = Vec::new();
4220
4221        // ── Prefill: forward each prompt token once, KEEP the last hidden.
4222        //    Dense prefill runs in fused pairs (weights streamed once per
4223        //    two positions — bit-identical to sequential, proven by the
4224        //    pair tests). With MTP: warm the draft head on
4225        //    (hidden_p, token_{p+1}) pairs.
4226        let mut hidden = vec![0.0f32; self.hidden_size];
4227        let mut pos = reuse_from;
4228        // lm_head-in-graph is only sound when the very next logits
4229        // consumer is this loop's own (MTP and skill routing interleave
4230        // other forwards / can swap lm_head between forward and sample).
4231        // CMF_GPU_LMHEAD=0 keeps lm_head off the graph: the token reads back
4232        // the 8 KB hidden instead of ~1 MB of logits, and the head runs on
4233        // the host. A probe for how much of the graph's fixed per-token cost
4234        // is the logits readback (the layer sweep puts that fixed part at
4235        // 3.88 ms of an 18.5 ms frame).
4236        let fuse_lm = mtp.is_none()
4237            && router.is_none()
4238            && std::env::var("CMF_GPU_LMHEAD").as_deref() != Ok("0");
4239        self.graph_logits = None;
4240        self.graph_want_logits = false;
4241        let _tpf = std::time::Instant::now();
4242        let batch_k = self.generation_batch_k();
4243        if let Some(rows) = prompt_rows {
4244            let hs = self.hidden_size;
4245            let chunk = self.prefill_chunk().max(1);
4246            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4247                let end = (pos + chunk).min(input_ids.len());
4248                let hb = match self.prefill_input_rows(
4249                    PrefillIn::Hidden(&rows[pos * hs..end * hs]), pos, task_mask,
4250                ) {
4251                    Ok(hb) => hb,
4252                    Err(err) => {
4253                        self.finish_generation(&mut mtp, &mut router, true);
4254                        return Err(err);
4255                    }
4256                };
4257                if mimo_spec { self.mimo_note_rows(&hb, pos); }
4258                hidden.copy_from_slice(&hb[hb.len() - hs..]);
4259                pos = end;
4260            }
4261        }
4262        // DeepSeek-V4 owns a separate hyper-connection stack. Route it
4263        // before the generic prefill choices: those correctly reject an
4264        // empty `weights.layers`, but their final per-position fallback used
4265        // to consume the whole prompt before `dsv4::forward_chunk` could see
4266        // it. The batch implementation therefore existed without a live
4267        // production entry point.
4268        //
4269        // Bounded chunks preserve cancellation responsiveness. Only the
4270        // prompt's final chunk asks for logits; every earlier head projection
4271        // would produce 129 280 values that no caller reads.
4272        while self.qwen4_exp.is_some()
4273            && mtp.is_none()
4274            && pos < input_ids.len()
4275            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4276        {
4277            // The device path takes several prompt tokens per layer frame;
4278            // the host path runs them one by one inside the same call.
4279            let end = (pos + crate::qwen4_exp::prefill_chunk()).min(input_ids.len());
4280            let want_logits = end == input_ids.len();
4281            let mut lg = Vec::new();
4282            if let Some(b) = &mut self.qwen4_exp {
4283                crate::qwen4_exp::forward_tokens(
4284                    &b.0,
4285                    &b.1,
4286                    &b.2,
4287                    &mut b.3,
4288                    &input_ids[pos..end],
4289                    pos,
4290                    &self.inv_freq,
4291                    self.pool.as_deref(),
4292                    &mut lg,
4293                    want_logits,
4294                );
4295            }
4296            if want_logits {
4297                self.graph_logits = Some(lg);
4298            }
4299            pos = end;
4300            hidden.fill(0.0);
4301        }
4302        while self.dsv4.is_some()
4303            && mtp.is_none()
4304            && pos < input_ids.len()
4305            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4306        {
4307            let end = (pos + prefill_chunk()).min(input_ids.len());
4308            let ids: Vec<u32> = input_ids[pos..end].to_vec();
4309            let mut lg = Vec::new();
4310            if let Some(b) = &mut self.dsv4 {
4311                let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
4312                crate::dsv4::forward_chunk(
4313                    g,
4314                    layers,
4315                    &cfg,
4316                    st,
4317                    &ids,
4318                    pos,
4319                    &self.inv_freq,
4320                    self.pool.as_deref(),
4321                    &mut lg,
4322                    end == input_ids.len(),
4323                );
4324            }
4325            if end == input_ids.len() {
4326                self.graph_logits = Some(lg);
4327            }
4328            pos = end;
4329            hidden = vec![0.0; self.hidden_size];
4330        }
4331        let dsv41_prefill = self.dsv41_prefill.take();
4332        while self.dsv41.is_some()
4333            && mtp.is_none()
4334            && pos < input_ids.len()
4335            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4336        {
4337            let end = (pos + prefill_chunk()).min(input_ids.len());
4338            let ids: Vec<u32> = input_ids[pos..end].to_vec();
4339            let mut lg = Vec::new();
4340            if let Some(b) = &mut self.dsv41 {
4341                let (g, layers, cfg, st) = (&b.0, &b.1, &b.2, &mut b.3);
4342                if let Some((embeddings, participates)) = dsv41_prefill.as_ref() {
4343                    crate::dsv41::forward_chunk_masked_with_embeddings(
4344                        g,
4345                        layers,
4346                        cfg,
4347                        st,
4348                        &ids,
4349                        pos,
4350                        &embeddings[pos..end],
4351                        &participates[pos..end],
4352                        self.pool.as_deref(),
4353                        &mut lg,
4354                    );
4355                } else {
4356                    crate::dsv41::forward_chunk(
4357                        g,
4358                        layers,
4359                        cfg,
4360                        st,
4361                        &ids,
4362                        pos,
4363                        self.pool.as_deref(),
4364                        &mut lg,
4365                    );
4366                }
4367            }
4368            if end == input_ids.len() {
4369                self.graph_logits = Some(lg);
4370            }
4371            pos = end;
4372            hidden = vec![0.0; self.hidden_size];
4373        }
4374        // With dynamic routing, prefill sequentially so the φ hook fires
4375        // over the PROMPT — the router enters decode with a warm φ (the
4376        // fused-pair path skips the per-layer φ capture). o1 layers
4377        // collect their query trace in both the single and pair paths.
4378        let dyn_prefill = router.is_some();
4379        // Optional bounded calibration prefix for generation.  The normal
4380        // O(1) path seals after the full prompt; this explicit knob instead
4381        // runs only the requested prefix through exact attention, seals the
4382        // Nyström state, and streams the rest of the prompt through the same
4383        // O(1) step used by decode.  It keeps the O(1) layers' Q trace and
4384        // temporary full KV bounded by the prefix while leaving the default
4385        // full-prompt quality profile untouched.
4386        let o1_prefill_limit = o1_prefill
4387            .and_then(|requested| self.o1_effective_boundary(requested))
4388            .map(|boundary| boundary.min(input_ids.len()));
4389        let mut o1_sealed = false;
4390        if let Some(limit) = o1_prefill_limit {
4391            // Reuse the exact batched prefix machinery when available; it
4392            // records the same per-position Q trace as the full prefill.
4393            if self.can_prefill_batched() && limit > 2 {
4394                let chunk = self.prefill_chunk();
4395                let hs = self.hidden_size;
4396                while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4397                    let end = (pos + chunk).min(limit);
4398                    let hb = self.prefill_batch(&input_ids[pos..end], pos);
4399                    hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4400                    pos = end;
4401                }
4402            } else {
4403                while pos < limit && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4404                    hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, None);
4405                    pos += 1;
4406                }
4407            }
4408            if pos >= limit {
4409                o1_sealed = match self.o1_seal_checked() {
4410                    Ok(sealed) => sealed,
4411                    Err(err) => {
4412                        self.finish_generation(&mut mtp, &mut router, true);
4413                        return Err(err);
4414                    }
4415                };
4416                tracing::info!(
4417                    "o1 bounded prompt prefix: requested={} effective={} processed={} of {} token(s)",
4418                    o1_prefill.unwrap_or(0),
4419                    self.o1_effective_boundary(o1_prefill.unwrap_or(0))
4420                        .unwrap_or(limit),
4421                    limit,
4422                    input_ids.len()
4423                );
4424            }
4425        }
4426        // q1 hybrids on Metal: the per-position GPU token graph beats
4427        // the CPU chunk-GEMM (whose wall is the sequential scalar GDN
4428        // recurrence), so prefill goes position-by-position through the
4429        // same graph as decode. Pure-attention models keep the batched
4430        // path — there the chunk-GEMM amortization wins.
4431        let graph_prefill = self.graph_prefill_preferred();
4432        // Native Metal, q4tp GDN hybrids: the prompt through the b-row
4433        // rows graph — projections as GEMMs over up to 512 positions, the
4434        // GDN recurrence in registers on the device, K/V rows appended by
4435        // the chunk — instead of one token-graph submit per position (the
4436        // 27B: 8 tok/s → GEMM-bound). The MTP warm-up rows come out of one
4437        // batched run of the block per chunk. Any refusal leaves the rest
4438        // of the prompt to the sequential paths below.
4439        #[cfg(target_os = "macos")]
4440        if task_mask.is_none()
4441            && !dyn_prefill
4442            && (crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in())
4443            && crate::gpu::enabled_here()
4444            && self.gdn_cfg.is_some()
4445            && self.g3n.is_none()
4446            && input_ids.len() > 8
4447            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err()
4448            && std::env::var("CMF_METAL_PREFILL").as_deref() != Ok("0")
4449        {
4450            let chunk: usize = std::env::var("CMF_METAL_PREFILL_CHUNK")
4451                .ok()
4452                .and_then(|v| v.parse().ok())
4453                .filter(|&v| (16..=512).contains(&v))
4454                .unwrap_or(256);
4455            let hs = self.hidden_size;
4456            let _tp = std::time::Instant::now();
4457            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4458                let end = (pos + chunk).min(input_ids.len());
4459                let hb = match self.prefill_batch_metal(&input_ids[pos..end], pos) {
4460                    MetalPrefillOutcome::Completed(hb) => hb,
4461                    MetalPrefillOutcome::Declined => break,
4462                    MetalPrefillOutcome::Failed => {
4463                        self.finish_generation(&mut mtp, &mut router, true);
4464                        return Err("ordinary Metal prefill failed after admission".into());
4465                    }
4466                };
4467                if let Some(m) = &mut mtp {
4468                    let n_pairs = if end < input_ids.len() {
4469                        end - pos
4470                    } else {
4471                        end - pos - 1
4472                    };
4473                    if n_pairs > 0 {
4474                        let pairs: Vec<(&[f32], u32)> = (0..n_pairs)
4475                            .map(|j| (&hb[j * hs..(j + 1) * hs], input_ids[pos + j + 1]))
4476                            .collect();
4477                        if !self.mtp_warm_batch_metal(m, &pairs, pos) {
4478                            for (j, (h, t)) in pairs.iter().enumerate() {
4479                                let h = h.to_vec();
4480                                let _ = self.mtp_step(m, &h, *t, pos + j);
4481                            }
4482                        }
4483                    }
4484                }
4485                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4486                pos = end;
4487            }
4488            if std::env::var("CMF_PREFILL_PROF").is_ok() {
4489                eprintln!(
4490                    "metal-prefill: {} of {} tokens in {:.1} ms",
4491                    pos,
4492                    input_ids.len(),
4493                    _tp.elapsed().as_secs_f64() * 1e3
4494                );
4495            }
4496        }
4497        self.mimo_moe_prepare();
4498        // A MoE stack larger than the card (MiMo-V2 q4tp on 96 GB): the
4499        // batched wgpu graph runs the device prefix of every chunk — its
4500        // experts resident — and the host's batched layer walk finishes
4501        // the chunk. Any refusal leaves the rest of the prompt to the
4502        // chunked prefill below.
4503        #[cfg(not(target_os = "macos"))]
4504        if task_mask.is_none()
4505            && !dyn_prefill
4506            && !graph_prefill
4507            && mtp.is_none()
4508            && o1_prefill.is_none()
4509            && !self.o1_active()
4510            && input_ids.len() > 2
4511            && self.batch_prefix_prefill()
4512        {
4513            let chunk = self.prefill_chunk().max(1);
4514            let hs = self.hidden_size;
4515            let t_bp = std::time::Instant::now();
4516            let pos0 = pos;
4517            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4518                let end = (pos + chunk).min(input_ids.len());
4519                let bk = end - pos;
4520                let mut hiddens = vec![0f32; bk * hs];
4521                for (j, &id) in input_ids[pos..end].iter().enumerate() {
4522                    hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4523                }
4524                let positions: Vec<usize> = (pos..end).collect();
4525                let mut run = 0usize;
4526                let outcome = self.try_batch_graph_wgpu_prefix(
4527                    &mut hiddens,
4528                    &positions,
4529                    bk,
4530                    None,
4531                    Some(&mut run),
4532                );
4533                match outcome {
4534                    crate::gpu::BatchGraphOutcome::Completed => {
4535                        let hb = if run < self.num_layers {
4536                            self.prefill_batch_span(
4537                                PrefillIn::Hidden(&hiddens),
4538                                pos,
4539                                None,
4540                                run,
4541                                self.num_layers,
4542                            )
4543                        } else {
4544                            hiddens
4545                        };
4546                        if mimo_spec {
4547                            self.mimo_note_rows(&hb, pos);
4548                        }
4549                        hidden.copy_from_slice(&hb[(bk - 1) * hs..]);
4550                        pos = end;
4551                    }
4552                    crate::gpu::BatchGraphOutcome::Failed => {
4553                        self.finish_generation(&mut mtp, &mut router, true);
4554                        return Err("batched prefix prefill failed after admission".into());
4555                    }
4556                    crate::gpu::BatchGraphOutcome::Declined => {
4557                        // Earlier chunks left their prefix rows on the
4558                        // device only: the host walk below needs them.
4559                        #[cfg(feature = "gpu")]
4560                        if pos > pos0 {
4561                            self.pull_lagging_host_kv(0, self.num_layers, pos);
4562                        }
4563                        break;
4564                    }
4565                }
4566            }
4567            if std::env::var("CMF_PREFILL_PROF").is_ok() {
4568                eprintln!(
4569                    "batch-prefix prefill: {} of {} tokens in {:.1} ms",
4570                    pos - pos0,
4571                    input_ids.len(),
4572                    t_bp.elapsed().as_secs_f64() * 1e3
4573                );
4574            }
4575        }
4576        if task_mask.is_none()
4577            && !dyn_prefill
4578            && !graph_prefill
4579            && self.can_prefill_batched()
4580            && self.g3n.is_none()
4581            && o1_prefill.is_none()
4582            && input_ids.len() > 2
4583        {
4584            // Production prefill = the same chunked prefill-GEMM that
4585            // bench/PPL measure (roadmap §3 P0: generation used to warm
4586            // the prompt with the slower pair path — the published
4587            // prefill number didn't match real TTFT). MTP warm-up reads
4588            // each position's hidden straight from the chunk result.
4589            let chunk = self.prefill_chunk();
4590            let hs = self.hidden_size;
4591            while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4592                let end = (pos + chunk).min(input_ids.len());
4593                let hb = self.prefill_batch(&input_ids[pos..end], pos);
4594                if mimo_spec {
4595                    self.mimo_note_rows(&hb, pos);
4596                }
4597                if let Some(m) = &mut mtp {
4598                    let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4599                        .ok()
4600                        .and_then(|v| v.parse().ok())
4601                        .unwrap_or(0);
4602                    for p in pos..end {
4603                        if p + 1 < input_ids.len() {
4604                            if probe >= 1 && p + 2 < input_ids.len() {
4605                                // Teacher-forced chain acceptance (see the
4606                                // tail loop's twin): the warm-up row stays,
4607                                // the chain's rows roll back.
4608                                let (d1, mut hx) = self.mtp_step_h(
4609                                    m,
4610                                    &hb[(p - pos) * hs..(p - pos + 1) * hs],
4611                                    input_ids[p + 1],
4612                                    p,
4613                                );
4614                                let mut ok = d1 == input_ids[p + 2];
4615                                Self::chain_probe_note(0, ok);
4616                                let mut d_prev = d1;
4617                                let mut extra = 0usize;
4618                                for j in 1..probe {
4619                                    if p + 2 + j >= input_ids.len() {
4620                                        break;
4621                                    }
4622                                    let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, p + 1 + j);
4623                                    extra += 1;
4624                                    ok = ok && dj == input_ids[p + 2 + j];
4625                                    Self::chain_probe_note(j, ok);
4626                                    d_prev = dj;
4627                                    hx = hj;
4628                                }
4629                                m.kv.truncate_last(extra);
4630                            } else {
4631                                let _ = self.mtp_step(
4632                                    m,
4633                                    &hb[(p - pos) * hs..(p - pos + 1) * hs],
4634                                    input_ids[p + 1],
4635                                    p,
4636                                );
4637                            }
4638                        }
4639                    }
4640                }
4641                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
4642                pos = end;
4643            }
4644        }
4645        let pair_off = std::env::var("CMF_PAIR").is_ok_and(|v| v == "0");
4646        if task_mask.is_none()
4647            && !dyn_prefill
4648            && !graph_prefill
4649            && !pair_off
4650            && self.pair_supported()
4651            && o1_prefill.is_none()
4652        {
4653            while pos + 1 < input_ids.len()
4654                && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4655            {
4656                let e1 = self.embed_single(input_ids[pos]);
4657                let e2 = self.embed_single(input_ids[pos + 1]);
4658                let (h1, h2) = self.forward_pair(&e1, &e2, pos);
4659                if mimo_spec {
4660                    self.mimo_note_rows(&h1, pos);
4661                    self.mimo_note_rows(&h2, pos + 1);
4662                }
4663                // Both prefill tokens are real → commit lane-2 states.
4664                self.commit_linear_scratch();
4665                if let Some(m) = &mut mtp {
4666                    let _ = self.mtp_step(m, &h1, input_ids[pos + 1], pos);
4667                    if pos + 2 < input_ids.len() {
4668                        let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4669                            .ok()
4670                            .and_then(|v| v.parse().ok())
4671                            .unwrap_or(0);
4672                        if probe >= 1 && pos + 3 < input_ids.len() {
4673                            // Same teacher-forced chain table as the tail
4674                            // loop below, fed from the pair path that owns
4675                            // most prefill positions.
4676                            let (d1, mut hx) = self.mtp_step_h(m, &h2, input_ids[pos + 2], pos + 1);
4677                            let mut ok = d1 == input_ids[pos + 3];
4678                            Self::chain_probe_note(0, ok);
4679                            let mut d_prev = d1;
4680                            let mut extra = 0usize;
4681                            for j in 1..probe {
4682                                if pos + 3 + j >= input_ids.len() {
4683                                    break;
4684                                }
4685                                let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 2 + j);
4686                                extra += 1;
4687                                ok = ok && dj == input_ids[pos + 3 + j];
4688                                Self::chain_probe_note(j, ok);
4689                                d_prev = dj;
4690                                hx = hj;
4691                            }
4692                            m.kv.truncate_last(extra);
4693                        } else {
4694                            let _ = self.mtp_step(m, &h2, input_ids[pos + 2], pos + 1);
4695                        }
4696                    }
4697                }
4698                hidden = h2;
4699                pos += 2;
4700            }
4701        }
4702        // Batched GPU prefill for the wgpu decode graph (GDN hybrids): K prompt
4703        // positions per submit — projections/FFN as GEMMs (weight once per K),
4704        // attention/GDN looped inside — instead of one whole-graph submit per
4705        // position. Falls through to the per-position graph on any refusal.
4706        // Batched prefill is opt-in (CMF_BATCH_K>0). Default 0 = per-position
4707        // graph prefill. (Steady-state decode is provably identical either way —
4708        // token-graph submit and lm_head both unchanged — so this only trades
4709        // prefill wall.)
4710        // A bounded O(1) prefix is the one post-seal prompt interval: only
4711        // admit its batch when the device O(1) route is explicitly enabled and
4712        // every sealed layer exposes a portable view. The same batch size and
4713        // refusal behavior remain the ordinary controls/comparator.
4714        let o1_batch_ready = o1_sealed
4715            && o1_prefill.is_some()
4716            && mtp.is_none()
4717            && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
4718            && (0..self.num_layers).all(|li| {
4719                let cache = &self.kv_cache.layers[self.phys_layer(li)];
4720                cache.o1.is_none() || cache.o1_views().is_some()
4721            });
4722        // The ordinary graph-prefill route can share each completed trunk
4723        // chunk with an attached MTP head.  Keep chain probing on its
4724        // established per-position path: the probe deliberately needs every
4725        // teacher-forced draft row and its rollback table.
4726        let mtp_batch_prefill = mtp.is_some()
4727            && graph_prefill
4728            && task_mask.is_none()
4729            && !dyn_prefill
4730            && !self.o1_active()
4731            && std::env::var("CMF_MTP_CHAIN_PROBE").is_err();
4732        if batch_k > 0
4733            && (graph_prefill || o1_batch_ready)
4734            && task_mask.is_none()
4735            && (!self.o1_active() || o1_batch_ready)
4736            && (mtp.is_none() || mtp_batch_prefill)
4737            && !dyn_prefill
4738            && pos + 1 < input_ids.len()
4739        {
4740            let hs = self.hidden_size;
4741            let chunk = batch_k;
4742            while pos < input_ids.len() {
4743                let end = (pos + chunk).min(input_ids.len());
4744                let bk = end - pos;
4745                let mut hiddens = vec![0f32; bk * hs];
4746                for (j, &id) in input_ids[pos..end].iter().enumerate() {
4747                    hiddens[j * hs..(j + 1) * hs].copy_from_slice(&self.embed_single(id));
4748                }
4749                let positions: Vec<usize> = (pos..end).collect();
4750                let t_chunk = std::time::Instant::now();
4751                let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
4752                let ok_b = outcome == crate::gpu::BatchGraphOutcome::Completed;
4753                if std::env::var("CMF_GRAPH_PROF").is_ok() {
4754                    let ms = t_chunk.elapsed().as_secs_f64() * 1000.0;
4755                    eprintln!(
4756                        "batch-chunk: phase=prompt mode={} k={bk} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
4757                        if o1_batch_ready {
4758                            "o1"
4759                        } else if mtp_batch_prefill {
4760                            "ordinary_mtp"
4761                        } else {
4762                            "ordinary"
4763                        },
4764                        bk as f64 / (ms / 1000.0)
4765                    );
4766                }
4767                {
4768                    use std::sync::atomic::{AtomicBool, Ordering};
4769                    static SAID: AtomicBool = AtomicBool::new(false);
4770                    if !SAID.swap(true, Ordering::Relaxed) {
4771                        if ok_b {
4772                            tracing::info!(
4773                                "batched prefill: ACTIVE mode={} (k={bk})",
4774                                if o1_batch_ready {
4775                                    "o1"
4776                                } else if mtp_batch_prefill {
4777                                    "ordinary_mtp"
4778                                } else {
4779                                    "ordinary"
4780                                }
4781                            );
4782                        } else {
4783                            tracing::warn!("batched prefill {:?} — per-position graph", outcome);
4784                        }
4785                    }
4786                }
4787                if ok_b {
4788                    if mimo_spec {
4789                        self.mimo_note_rows(&hiddens, pos);
4790                    }
4791                    if mtp_batch_prefill {
4792                        let n_pairs = mtp_prefill_pair_count(pos, end, input_ids.len());
4793                        if n_pairs > 0 {
4794                            // `hiddens` is owned by this chunk, so materialize
4795                            // row slices before borrowing the detached MTP
4796                            // module.  The last prompt row has no successor;
4797                            // the helper above is the single source of that
4798                            // boundary rule.
4799                            let rows: Vec<Vec<f32>> = (0..n_pairs)
4800                                .map(|j| hiddens[j * hs..(j + 1) * hs].to_vec())
4801                                .collect();
4802                            let pairs: Vec<(&[f32], u32)> = rows
4803                                .iter()
4804                                .enumerate()
4805                                .map(|(j, row)| (row.as_slice(), input_ids[pos + j + 1]))
4806                                .collect();
4807                            if std::env::var("CMF_GRAPH_PROF").is_ok() {
4808                                eprintln!(
4809                                    "mtp-warm: phase=prompt mode=ordinary_mtp first_pos={} pairs={} last_pos={}",
4810                                    pos,
4811                                    n_pairs,
4812                                    pos + n_pairs - 1,
4813                                );
4814                            }
4815                            let warm_error = if let Some(m) = mtp.as_mut() {
4816                                self.mtp_warm_prefill_pairs(m, &pairs, pos).err()
4817                            } else {
4818                                None
4819                            };
4820                            if let Some(err) = warm_error {
4821                                // The trunk batch was already admitted.  A
4822                                // failed MTP warm-up therefore clears both
4823                                // mirrors and exits; continuing would pair a
4824                                // current trunk state with a stale MTP cache.
4825                                self.finish_generation(&mut mtp, &mut router, true);
4826                                return Err(err.to_string());
4827                            }
4828                        }
4829                    }
4830                    hidden.copy_from_slice(&hiddens[(bk - 1) * hs..]);
4831                    pos = end;
4832                } else if outcome == crate::gpu::BatchGraphOutcome::Failed {
4833                    // A failed batch may have advanced a device recurrent
4834                    // state (ordinary GDN or sealed O(1)). A CPU fallback
4835                    // would then observe stale accumulators, so clear the
4836                    // request state and make the failure explicit.
4837                    self.finish_generation(&mut mtp, &mut router, true);
4838                    return Err(if o1_batch_ready {
4839                        "sealed O(1) batch graph failed after admission".to_string()
4840                    } else {
4841                        "ordinary recurrent batch graph failed after admission".to_string()
4842                    });
4843                } else {
4844                    break; // unsupported → per-position graph handles the rest
4845                }
4846            }
4847        }
4848        // Resident Embryo graph: the prompt in chunks of one submit each
4849        // instead of one whole-graph submit per position; the last chunk
4850        // carries the logits exactly as the per-position walk would.
4851        if graph_prefill
4852            && task_mask.is_none()
4853            && mtp.is_none()
4854            && !dyn_prefill
4855            && pos == 0
4856            && input_ids.len() > 1
4857            && !self.cancel.load(std::sync::atomic::Ordering::Relaxed)
4858        {
4859            if let Some(lg) = self.embryo_prefill_chunked(input_ids, 0) {
4860                self.graph_logits = Some(lg);
4861                hidden = vec![0.0; self.hidden_size];
4862                pos = input_ids.len();
4863            }
4864        }
4865        while pos < input_ids.len() && !self.cancel.load(std::sync::atomic::Ordering::Relaxed) {
4866            self.graph_want_logits = fuse_lm && pos + 1 == input_ids.len();
4867            hidden = self.forward_layers(&self.embed_single(input_ids[pos]), pos, task_mask);
4868            if mimo_spec {
4869                self.mimo_note_rows(&hidden, pos);
4870            }
4871            if let Some(m) = &mut mtp {
4872                if pos + 1 < input_ids.len() {
4873                    // `CMF_MTP_CHAIN_PROBE=k`: teacher-forced acceptance of a
4874                    // CHAINED draft — iterate the head on its own hidden k
4875                    // deep and score every depth against the prompt's real
4876                    // continuation. The economics of a k-token speculative
4877                    // round stand or fall on this table.
4878                    let probe: usize = std::env::var("CMF_MTP_CHAIN_PROBE")
4879                        .ok()
4880                        .and_then(|v| v.parse().ok())
4881                        .unwrap_or(0);
4882                    if probe >= 1 && pos + 2 < input_ids.len() {
4883                        let (d1, mut hx) = self.mtp_step_h(m, &hidden, input_ids[pos + 1], pos);
4884                        let mut ok = d1 == input_ids[pos + 2];
4885                        Self::chain_probe_note(0, ok);
4886                        let mut d_prev = d1;
4887                        let mut extra = 0usize;
4888                        for j in 1..probe {
4889                            if pos + 2 + j >= input_ids.len() {
4890                                break;
4891                            }
4892                            let (dj, hj) = self.mtp_step_h(m, &hx, d_prev, pos + 1 + j);
4893                            extra += 1;
4894                            ok = ok && dj == input_ids[pos + 2 + j];
4895                            Self::chain_probe_note(j, ok);
4896                            d_prev = dj;
4897                            hx = hj;
4898                        }
4899                        // The chain's rows are speculation, not the prompt —
4900                        // keep only the warmup row the plain path would add.
4901                        m.kv.truncate_last(extra);
4902                    } else {
4903                        let _ = self.mtp_step(m, &hidden, input_ids[pos + 1], pos);
4904                    }
4905                }
4906            }
4907            pos += 1;
4908        }
4909        if std::env::var("CMF_PREFILL_PROF").is_ok() {
4910            eprintln!(
4911                "prefill: {} tokens in {:.1} ms (batch_k={batch_k})",
4912                input_ids.len(),
4913                _tpf.elapsed().as_secs_f64() * 1000.0
4914            );
4915        }
4916        if self
4917            .graph_failed
4918            .swap(false, std::sync::atomic::Ordering::Relaxed)
4919        {
4920            // MTP is detached for speculative generation.  Restore the
4921            // module before returning the terminal graph error; otherwise a
4922            // failed request would silently remove the head from a pooled
4923            // pipeline and the next request would lose its configured route.
4924            self.finish_generation(&mut mtp, &mut router, true);
4925            return Err("GPU token graph failed during prefill".to_string());
4926        }
4927        // Cancelled mid-prefill: the cache holds a partial prompt —
4928        // drop the reuse history and return an empty generation.
4929        if self
4930            .cancel
4931            .swap(false, std::sync::atomic::Ordering::Relaxed)
4932        {
4933            // A cancelled prefill can already have advanced the device
4934            // mirror. Drop the whole partial sequence so a pooled pipeline
4935            // cannot carry that state into its next request.
4936            self.finish_generation(&mut mtp, &mut router, true);
4937            return Ok(GenerateResult {
4938                text: String::new(),
4939                token_ids: Vec::new(),
4940                prompt_tokens: input_ids.len(),
4941                tokens_generated: 0,
4942                finish_reason: "cancelled".to_string(),
4943                mtp_drafted: 0,
4944                mtp_accepted: 0,
4945                token_confidence: Vec::new(),
4946                traces: Vec::new(),
4947            });
4948        }
4949
4950        // Prompt absorbed → freeze the o1 layers' skeletons; from here
4951        // every decode step on those layers is O(W + m·dv + m²).
4952        if !o1_sealed {
4953            match self.o1_seal_checked() {
4954                Ok(_) => {}
4955                Err(err) => {
4956                    self.finish_generation(&mut mtp, &mut router, true);
4957                    return Err(err);
4958                }
4959            }
4960        }
4961
4962        // Commit one token: push, check EOS, stream. Returns false = stop.
4963        macro_rules! commit {
4964            ($id:expr) => {{
4965                all_ids.push($id);
4966                generated += 1;
4967                self.note_draft_id($id);
4968                if self.tokenizer.is_eos($id) && !self.ignore_eos {
4969                    finish_reason = "stop".to_string();
4970                    false
4971                } else {
4972                    let token_text = self.tokenizer.decode_token($id);
4973                    let mut go = true;
4974                    if let Some(ref mut cb) = on_token {
4975                        if !cb(&token_text) {
4976                            finish_reason = "cancelled".to_string();
4977                            go = false;
4978                        }
4979                    }
4980                    go
4981                }
4982            }};
4983        }
4984
4985        // Speculation is decided by MEASUREMENT, not by an acceptance
4986        // model. A k=4 round costs ~3.8 plain tokens on the 5090 (draft
4987        // 6.6 + verify 66.6 + commit 4.8 ms against a 20.6 ms token), so it
4988        // pays only when the head lands ~2.8 of 4 — predictable text (code,
4989        // structured output) does, free prose often does not, and the
4990        // ratio at which the two cross depends on the card and the context
4991        // depth. So: four speculative rounds timed, then eight plain
4992        // tokens timed, and the faster arm runs until a re-check 256
4993        // tokens later (context growth moves the balance). The trial
4994        // costs at most a few tokens of the slower arm per 256.
4995        let mut spec_trial = SpecTrial::Spec {
4996            t0: std::time::Instant::now(),
4997            gen0: generated,
4998            rounds: 0,
4999        };
5000        // The token-count proxy prices a round at ~1.9 plain tokens. That
5001        // holds for the Metal rounds whose cost was measured — greedy and
5002        // the sparse sampling chain — so an expensive round (the dense
5003        // chain, reachable only by `CMF_GRAPH_SPEC_SAMPLE=1`) still times
5004        // the plain path before it decides.
5005        let mut spec_mon = SpecMon {
5006            metal: graph_spec && crate::gpu::q1_force() && spec_cheap_round,
5007            ..SpecMon::default()
5008        };
5009        let mut spec_watchdog_off = false;
5010        // CMF_GRAPH_SPEC_TIME: the round walls so far (round 1 excluded —
5011        // it pays the scratch), for the outlier test on each new one
5012        let mut spec_walls: Vec<f32> = Vec::new();
5013        // ... and the end of the last round: the host time between rounds
5014        // (token commits, streaming, the loop top) is printed at level 2
5015        let mut spec_round_end: Option<std::time::Instant> = None;
5016        if mimo_spec {
5017            if let Ok(path) = std::env::var("CMF_MIMO_MTP_PROBE") {
5018                if let Some(mut st) = self.mimo_mtp.take() {
5019                    self.mimo_mtp_probe(&mut st, input_ids, &path);
5020                    self.mimo_mtp = Some(st);
5021                }
5022            }
5023        }
5024        // ── Decode ──
5025        let mut next_pos = input_ids.len();
5026        'decode: while generated < max_tokens {
5027            if self
5028                .graph_failed
5029                .swap(false, std::sync::atomic::Ordering::Relaxed)
5030            {
5031                // Keep the detached MTP module attached after a terminal
5032                // graph error so the pipeline can be reused for a fresh
5033                // sequence.  `clear_sequence_state` only clears mirrors and
5034                // host KV; it cannot recover a module dropped here.
5035                self.finish_generation(&mut mtp, &mut router, true);
5036                return Err("GPU token graph failed during decode".to_string());
5037            }
5038            if self
5039                .cancel
5040                .swap(false, std::sync::atomic::Ordering::Relaxed)
5041            {
5042                finish_reason = "cancelled".to_string();
5043                break 'decode;
5044            }
5045            // A rejected speculative draft already drew this position's
5046            // token from the residual distribution (graph_spec_step); it
5047            // is committed as-is — sampling again from the row's logits
5048            // would bias the stream toward the target's mode.
5049            if mimo_spec && next_pos > 0 {
5050                // Every path leaves `hidden` = the backbone output at
5051                // next_pos-1; the draft layers read it (idempotent).
5052                self.mimo_note_rows(&hidden, next_pos - 1);
5053            }
5054            let forced = self.spec_forced.take();
5055            let mut logits = match (forced, self.graph_logits.take()) {
5056                (Some(_), _) => Vec::new(),
5057                (None, Some(lg)) => lg,
5058                (None, None) => {
5059                    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Head);
5060                    inference::rms_norm_into(
5061                        &hidden,
5062                        &self.weights.final_norm,
5063                        self.rms_eps,
5064                        self.norm_style,
5065                        &mut self.ws.n1,
5066                    );
5067                    self.lm_head_forward(&self.ws.n1)
5068                }
5069            };
5070            // CMF_LOGIT_DUMP=<path>: the first decode step's hidden + logits
5071            // as raw f32 (hidden first) — cross-backend numerics diffing.
5072            if generated
5073                == std::env::var("CMF_LOGIT_DUMP_STEP")
5074                    .ok()
5075                    .and_then(|v| v.parse().ok())
5076                    .unwrap_or(0)
5077            {
5078                if let Ok(path) = std::env::var("CMF_LOGIT_DUMP") {
5079                    let mut bytes: Vec<u8> = Vec::with_capacity((hidden.len() + logits.len()) * 4);
5080                    for v in hidden.iter().chain(logits.iter()) {
5081                        bytes.extend_from_slice(&v.to_le_bytes());
5082                    }
5083                    if let Err(e) = std::fs::write(&path, &bytes) {
5084                        eprintln!("logit dump: failed to write {path}: {e}");
5085                        self.finish_generation(&mut mtp, &mut router, true);
5086                        return Err(format!("logit dump write failed: {e}"));
5087                    }
5088                }
5089            }
5090            // CMF_LOGIT_DUMP_ALL=<dir>: every decode step's logits as raw
5091            // f32, `<dir>/step{n:05}.f32` — step-by-step backend diffing
5092            // (a greedy run on two backends compares until they diverge).
5093            if let Ok(dir) = std::env::var("CMF_LOGIT_DUMP_ALL") {
5094                if !logits.is_empty() {
5095                    let path = std::path::Path::new(&dir).join(format!("step{generated:05}.f32"));
5096                    let bytes: Vec<u8> = logits.iter().flat_map(|v| v.to_le_bytes()).collect();
5097                    if let Err(e) =
5098                        std::fs::create_dir_all(&dir).and_then(|_| std::fs::write(&path, &bytes))
5099                    {
5100                        eprintln!("logit dump: failed to write {}: {e}", path.display());
5101                    }
5102                }
5103            }
5104            let t_next = match forced {
5105                Some(c) => c,
5106                None => {
5107                    let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Sampler);
5108                    sampler::sample_with_scratch_pool(
5109                        &logits,
5110                        &self.sampler_config,
5111                        self.sampler_config.penalty_past(&all_ids, bounded_native),
5112                        &mut self.rng,
5113                        &mut self.sampler_scratch,
5114                        self.pool.as_deref(),
5115                    )
5116                }
5117            };
5118            if self.confidence_on {
5119                confidence.push(if logits.is_empty() {
5120                    0.0
5121                } else {
5122                    sampler::top1_prob_pool(
5123                        self.pool.as_deref(),
5124                        &mut self.sampler_scratch,
5125                        &logits,
5126                        t_next,
5127                        calib_temp,
5128                    )
5129                });
5130            }
5131            if !logits.is_empty() {
5132                attention::recycle_buf(&mut logits);
5133            }
5134            if trace_on {
5135                // active_skill = the overlay in force while this token was
5136                // generated; recon/switched are filled after the post-emit
5137                // routing eval below (freshest coherence for this token).
5138                let skill = router.as_ref().and_then(|r| r.active_id());
5139                traces.push(TokenTrace {
5140                    t: generated,
5141                    token_id: t_next,
5142                    confidence: confidence.last().copied().unwrap_or(0.0),
5143                    active_skill: skill,
5144                    recon: None,
5145                    switched: false,
5146                });
5147            }
5148            if !commit!(t_next) {
5149                break 'decode;
5150            }
5151            if generated >= max_tokens {
5152                break 'decode;
5153            }
5154
5155            // Catch-all for decode paths that bypass the walks' own trim
5156            // (a graph step, a speculative round).
5157            self.swa_trim_tails();
5158            if self.dsv41.is_none() && self.kv_cache.needs_eviction() {
5159                // Say it ONCE, loudly: past this point the model keeps
5160                // talking but has lost half its context, and on a GDN
5161                // hybrid the graph's device state goes stale on top. The
5162                // Qwen3.8 bring-up spent a day reading this cliff as
5163                // three different model bugs.
5164                static SAID: std::sync::Once = std::sync::Once::new();
5165                SAID.call_once(|| {
5166                    tracing::warn!(
5167                        "KV cache full at {} positions — evicting half; quality \
5168                         will degrade. Raise CMF_MAX_SEQ.",
5169                        self.kv_cache.max_seq_len,
5170                    );
5171                });
5172                let keep = (self.kv_cache.max_seq_len / 2).max(1);
5173                self.kv_cache.evict(keep);
5174            }
5175
5176            // Advance the speculation trial: plain-phase accounting and
5177            // the periodic re-check happen here, on every token.
5178            if graph_spec {
5179                match spec_trial {
5180                    SpecTrial::Plain { t0, gen0 } if spec_mon.plain_done(t0, gen0, generated) => {
5181                        spec_mon.plain_ms =
5182                            t0.elapsed().as_secs_f64() * 1e3 / (generated - gen0) as f64;
5183                        let keep = spec_mon.pays();
5184                        tracing::info!(
5185                            "speculation trial: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5186                            spec_mon.tokens,
5187                            spec_mon.round_ms,
5188                            spec_mon.plain_ms,
5189                            if keep { "speculating" } else { "plain" }
5190                        );
5191                        spec_mon.fails = 0;
5192                        spec_trial = SpecTrial::Decided {
5193                            spec: keep,
5194                            recheck_at: if keep { usize::MAX } else { generated + 128 },
5195                        };
5196                    }
5197                    SpecTrial::Decided { recheck_at, .. } if generated >= recheck_at => {
5198                        spec_mon.n = 0;
5199                        spec_trial = SpecTrial::Spec {
5200                            t0: std::time::Instant::now(),
5201                            gen0: generated,
5202                            rounds: 0,
5203                        };
5204                    }
5205                    _ => {}
5206                }
5207                spec_watchdog_off = matches!(
5208                    spec_trial,
5209                    SpecTrial::Plain { .. } | SpecTrial::Decided { spec: false, .. }
5210                );
5211            }
5212            // ── MiMo-V2 draft stack: draft K, verify K+1 rows in one batch ──
5213            if mimo_spec && generated + 1 < max_tokens && next_pos > 0 {
5214                let budget = max_tokens - generated - 1;
5215                if let Some(mut st) = self.mimo_mtp.take() {
5216                    let k = st.depth.min(budget);
5217                    let r = self.mimo_spec_round(&mut st, next_pos, &all_ids, k);
5218                    self.mimo_mtp = Some(st);
5219                    let r = match r {
5220                        Ok(r) => r,
5221                        Err(err) => {
5222                            self.finish_generation(&mut mtp, &mut router, true);
5223                            return Err(err);
5224                        }
5225                    };
5226                    if let Some(r) = r {
5227                        drafted += r.drafted;
5228                        accepted += r.accepted.len();
5229                        let mut stopped = false;
5230                        for &id in &r.accepted {
5231                            if self.confidence_on {
5232                                confidence.push(0.0);
5233                            }
5234                            if !commit!(id) {
5235                                stopped = true;
5236                                break;
5237                            }
5238                        }
5239                        if stopped {
5240                            break 'decode;
5241                        }
5242                        next_pos += r.accepted.len() + 1;
5243                        hidden = r.hidden;
5244                        // The loop top chooses the round's own token from
5245                        // these logits — the same sampler, same history.
5246                        self.graph_logits = Some(r.logits);
5247                        continue 'decode;
5248                    }
5249                }
5250            }
5251            // ── Qwen3.8-Flash-Next draft head: k greedy drafts from the MTP
5252            //    sidecar, one batched verify window on the device ──
5253            #[cfg(feature = "gpu")]
5254            if self.speculative
5255                && self.qwen4_exp.is_some()
5256                && task_mask.is_none()
5257                && self.sampler_config.temperature < 1e-6
5258                && generated + 1 < max_tokens
5259                && next_pos > 0
5260                && std::env::var("CMF_QWEN_MTP").as_deref() != Ok("0")
5261            {
5262                let r = match &mut self.qwen4_exp {
5263                    Some(b) => crate::qwen4_exp::spec_round(
5264                        &b.0,
5265                        &b.1,
5266                        &b.2,
5267                        &mut b.3,
5268                        next_pos,
5269                        &all_ids,
5270                        &self.inv_freq,
5271                        self.pool.as_deref(),
5272                    ),
5273                    None => None,
5274                };
5275                if let Some(r) = r {
5276                    drafted += r.drafted;
5277                    accepted += r.accepted.len();
5278                    let mut stopped = false;
5279                    for &id in &r.accepted {
5280                        if self.confidence_on {
5281                            confidence.push(0.0);
5282                        }
5283                        if !commit!(id) {
5284                            stopped = true;
5285                            break;
5286                        }
5287                    }
5288                    if stopped {
5289                        break 'decode;
5290                    }
5291                    next_pos += r.accepted.len() + 1;
5292                    hidden.fill(0.0);
5293                    self.graph_logits = Some(r.logits);
5294                    continue 'decode;
5295                }
5296            }
5297            match &mut mtp {
5298                // ── Graph speculation: chain-draft, batch-verify on device ──
5299                #[cfg(feature = "gpu")]
5300                Some(m)
5301                    if graph_spec
5302                        && !spec_watchdog_off
5303                        && generated + 1 < max_tokens
5304                        && next_pos > 0 =>
5305                {
5306                    let t_round = std::time::Instant::now();
5307                    if spec_time_level() >= 2 {
5308                        if let Some(t) = spec_round_end.take() {
5309                            eprintln!(
5310                                "spec-gap {:.2} ms (host between rounds)",
5311                                t.elapsed().as_secs_f64() * 1e3
5312                            );
5313                        }
5314                    }
5315                    spec_stamps_begin();
5316                    // device buffers allocated during this round: a
5317                    // first-touch Shared allocation is zero-filled inside
5318                    // the command buffer that uses it, which is what the
5319                    // long outlier rounds were
5320                    #[cfg(target_os = "macos")]
5321                    let allocs0 = crate::gpu_metal::IO_BUF_ALLOCS
5322                        .load(std::sync::atomic::Ordering::Relaxed);
5323                    #[cfg(not(target_os = "macos"))]
5324                    let allocs0 = 0u64;
5325                    if let Some((extra, n_pos, new_h)) = self.graph_spec_step(
5326                        m,
5327                        &hidden,
5328                        t_next,
5329                        next_pos,
5330                        &mut drafted,
5331                        &mut accepted,
5332                        &mut all_ids,
5333                        max_tokens - generated,
5334                    ) {
5335                        next_pos = n_pos;
5336                        hidden = new_h;
5337                        let level = spec_time_level();
5338                        if level > 0 {
5339                            let wall = t_round.elapsed().as_secs_f32() * 1e3;
5340                            let stamps = spec_stamps_take();
5341                            // the running median of the rounds before this
5342                            // one (round 1 pays the scratch: not a sample)
5343                            let median = if spec_walls.len() >= 3 {
5344                                let mut s = spec_walls.clone();
5345                                s.sort_by(|a, b| a.partial_cmp(b).unwrap());
5346                                Some(s[s.len() / 2])
5347                            } else {
5348                                None
5349                            };
5350                            let outlier = median.is_some_and(|m| wall > 1.4 * m);
5351                            #[cfg(target_os = "macos")]
5352                            let allocs = crate::gpu_metal::IO_BUF_ALLOCS
5353                                .load(std::sync::atomic::Ordering::Relaxed)
5354                                - allocs0;
5355                            #[cfg(not(target_os = "macos"))]
5356                            let allocs = allocs0;
5357                            eprintln!(
5358                                "spec-round wall {wall:.1} ms → {} tokens{}{}",
5359                                extra.len() + 1,
5360                                if allocs > 0 {
5361                                    format!(" [{allocs} new device buffers]")
5362                                } else {
5363                                    String::new()
5364                                },
5365                                match (outlier, median) {
5366                                    (true, Some(m)) => format!(" OUTLIER (median {m:.1})"),
5367                                    _ => String::new(),
5368                                }
5369                            );
5370                            if level >= 2 || outlier {
5371                                let sum: f32 = stamps.iter().map(|s| s.1).sum();
5372                                eprintln!(
5373                                    "spec-stamps: {}| untracked {:.1}",
5374                                    spec_stamps_format(&stamps),
5375                                    wall - sum
5376                                );
5377                            }
5378                            if spec_mon.n >= 1 {
5379                                spec_walls.push(wall);
5380                            }
5381                        }
5382                        // One speculative round done: the monitor counts it
5383                        // (round 1 untimed — it pays the batch scratch and
5384                        // the draft mirror), and the trial advances.
5385                        spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, extra.len() + 1);
5386                        // the round's tokens land in `generated` below; the
5387                        // plain phase must start counting AFTER them
5388                        spec_trial = Self::spec_trial_round(
5389                            spec_trial,
5390                            &mut spec_mon,
5391                            generated + extra.len() + 1,
5392                        );
5393                        let mut stopped = false;
5394                        for &id in &extra {
5395                            if self.confidence_on {
5396                                confidence.push(0.0);
5397                            }
5398                            if !commit!(id) {
5399                                stopped = true;
5400                                break;
5401                            }
5402                        }
5403                        if stopped {
5404                            break 'decode;
5405                        }
5406                        if spec_time_level() >= 2 {
5407                            spec_round_end = Some(std::time::Instant::now());
5408                        }
5409                        continue 'decode;
5410                    }
5411                    if self
5412                        .graph_failed
5413                        .swap(false, std::sync::atomic::Ordering::Relaxed)
5414                    {
5415                        // `graph_spec_step` may have detached MTP while a
5416                        // warm-up was in flight.  Do not reinterpret its
5417                        // terminal device failure as a plain decode step;
5418                        // restore the head, clear both mirrors, and surface
5419                        // one explicit error to the caller.
5420                        self.finish_generation(&mut mtp, &mut router, true);
5421                        return Err("GPU MTP graph failed during speculative decode".to_string());
5422                    }
5423                    // Declined (batch graph refused): plain forward below —
5424                    // and a round that produced one token for the trial's
5425                    // ledger, so a graph that keeps refusing is measured out
5426                    // like a head that keeps missing (it was spinning
5427                    // forever on a file whose batch graph declines).
5428                    // A declined round is not a cheap one-token round — it
5429                    // is a verify that does not exist for this file (a
5430                    // healed q8_2f tail measured 760 drafts, 0 accepted, 33
5431                    // against 48.8 tok/s while the monitor called the draft
5432                    // alone "paying"). Count it as the losing streak in one.
5433                    spec_mon.round(t_round.elapsed().as_secs_f64() * 1e3, 1);
5434                    spec_mon.tokens = 0.0;
5435                    spec_mon.fails = 3;
5436                    spec_trial = Self::spec_trial_round(spec_trial, &mut spec_mon, generated + 1);
5437                    hidden = self.forward_layers(&self.embed_single(t_next), next_pos, task_mask);
5438                    next_pos += 1;
5439                    continue 'decode;
5440                }
5441                // ── Speculative: draft t+2, verify in a fused pair ──
5442                Some(m) if !graph_spec && generated + 1 < max_tokens => {
5443                    let draft = self.mtp_step(m, &hidden, t_next, next_pos - 1);
5444                    drafted += 1;
5445                    let emb1 = self.embed_single(t_next);
5446                    let emb2 = self.embed_single(draft);
5447                    let (h1, h2) = self.forward_pair(&emb1, &emb2, next_pos);
5448
5449                    inference::rms_norm_into(
5450                        &h1,
5451                        &self.weights.final_norm,
5452                        self.rms_eps,
5453                        self.norm_style,
5454                        &mut self.ws.n1,
5455                    );
5456                    let mut logits1 = self.lm_head_forward(&self.ws.n1);
5457                    let t_after = sampler::sample_with_scratch_pool(
5458                        &logits1,
5459                        &self.sampler_config,
5460                        self.sampler_config.penalty_past(&all_ids, bounded_native),
5461                        &mut self.rng,
5462                        &mut self.sampler_scratch,
5463                        self.pool.as_deref(),
5464                    );
5465                    if self.confidence_on {
5466                        confidence.push(sampler::top1_prob_pool(
5467                            self.pool.as_deref(),
5468                            &mut self.sampler_scratch,
5469                            &logits1,
5470                            t_after,
5471                            calib_temp,
5472                        ));
5473                    }
5474                    attention::recycle_buf(&mut logits1);
5475                    if trace_on {
5476                        // Speculative decode is mutually exclusive with
5477                        // dynamic routing (router is None here) — no skill.
5478                        traces.push(TokenTrace {
5479                            t: generated,
5480                            token_id: t_after,
5481                            confidence: confidence.last().copied().unwrap_or(0.0),
5482                            active_skill: None,
5483                            recon: None,
5484                            switched: false,
5485                        });
5486                    }
5487                    let stop = !commit!(t_after);
5488
5489                    if t_after == draft {
5490                        accepted += 1;
5491                        self.commit_linear_scratch();
5492                        let _ = self.mtp_step(m, &h1, t_after, next_pos);
5493                        hidden = h2;
5494                        next_pos += 2;
5495                    } else {
5496                        // The draft lane is wrong: roll its KV entry back.
5497                        for layer in &mut self.kv_cache.layers {
5498                            layer.truncate_last(1);
5499                        }
5500                        if !stop {
5501                            let _ = self.mtp_step(m, &h1, t_after, next_pos);
5502                            hidden = self.forward_layers(
5503                                &self.embed_single(t_after),
5504                                next_pos + 1,
5505                                None,
5506                            );
5507                        }
5508                        next_pos += 2;
5509                    }
5510                    if stop {
5511                        break 'decode;
5512                    }
5513                }
5514                // ── Vanilla: forward the sampled token ──
5515                _ => {
5516                    // ── DeepSeek-V4 speculative decode (CMF_DSV4_SPEC=1):
5517                    // draft five on the card, verify batched, commit the
5518                    // accepted prefix. Greedy only; a rejected token's state
5519                    // is restored and replayed, so output equals the walk. ──
5520                    #[cfg(feature = "gpu")]
5521                    if Self::dsv4_spec_on() && self.dsv4.is_some() {
5522                        static SAID: std::sync::Once = std::sync::Once::new();
5523                        SAID.call_once(|| {
5524                            eprintln!(
5525                                "dsv4-spec гейт: mtp={} mask={} router={} trace={} temp={} rep={} ",
5526                                !self.dsv4_mtp.is_empty(),
5527                                task_mask.is_none(),
5528                                router.is_none(),
5529                                !trace_on,
5530                                self.sampler_config.temperature < 1e-6,
5531                                self.sampler_config.repetition_penalty == 1.0,
5532                            );
5533                        });
5534                    }
5535                    #[cfg(feature = "gpu")]
5536                    if Self::dsv4_spec_on()
5537                        && self.dsv4.is_some()
5538                        && !self.dsv4_mtp.is_empty()
5539                        && task_mask.is_none()
5540                        && router.is_none()
5541                        && !trace_on
5542                        && self.sampler_config.temperature < 1e-6
5543                        && self.sampler_config.repetition_penalty == 1.0
5544                        && generated + 1 < max_tokens
5545                        && all_ids.len() >= 2
5546                        && generated >= dsv4_spec_retry_at
5547                    {
5548                        let tip_token = all_ids[all_ids.len() - 2];
5549                        let drafted0 = drafted;
5550                        let round = self.dsv4_spec_step(
5551                            tip_token,
5552                            t_next,
5553                            next_pos,
5554                            max_tokens.saturating_sub(generated),
5555                            &mut drafted,
5556                            &mut accepted,
5557                        );
5558                        if drafted > drafted0 {
5559                            let useful = round.as_ref().is_some_and(|(extra, _)| !extra.is_empty());
5560                            if useful {
5561                                dsv4_spec_bad = 0;
5562                            } else {
5563                                dsv4_spec_bad += 1;
5564                                if dsv4_spec_bad >= 2 {
5565                                    dsv4_spec_bad = 0;
5566                                    dsv4_spec_retry_at = generated.saturating_add(32);
5567                                    tracing::info!(
5568                                        "dsv4: draft не окупился дважды — точный walk на 32 токена"
5569                                    );
5570                                }
5571                            }
5572                        }
5573                        if let Some((extra, n_pos)) = round {
5574                            next_pos = n_pos;
5575                            let mut stopped = false;
5576                            for &id in &extra {
5577                                if self.confidence_on {
5578                                    confidence.push(0.0);
5579                                }
5580                                if !commit!(id) {
5581                                    stopped = true;
5582                                    break;
5583                                }
5584                            }
5585                            if stopped {
5586                                break 'decode;
5587                            }
5588                            continue 'decode;
5589                        }
5590                    }
5591                    self.graph_want_logits = fuse_lm;
5592                    // Greedy burst (CMF_MULTISTEP, default 8, 1 = off): while
5593                    // nothing observes per-token state — pure argmax sampling,
5594                    // no router/trace/confidence/mask — decode k tokens per
5595                    // submit and commit them wholesale. The trailing normal
5596                    // forward leaves logits for the loop top, as always.
5597                    let mut t_fwd = t_next;
5598                    let pure_greedy = self.sampler_config.temperature < 1e-6
5599                        && self.sampler_config.repetition_penalty == 1.0
5600                        && self.sampler_config.suppress_tokens.is_empty();
5601                    // Off by default: at every k the burst measured at or
5602                    // below the plain path on this graph shape (k=1 loses
5603                    // the argmax dispatches vs a 1 MB readback, k>=8 loses
5604                    // inter-step drains vs the saved sync). Experimental.
5605                    let burst_k = std::env::var("CMF_MULTISTEP")
5606                        .ok()
5607                        .and_then(|v| v.parse::<usize>().ok())
5608                        .unwrap_or(0);
5609                    if pure_greedy
5610                        && burst_k >= 1
5611                        && fuse_lm
5612                        && task_mask.is_none()
5613                        && router.is_none()
5614                        && !trace_on
5615                        && !self.confidence_on
5616                    {
5617                        let mut stopped = false;
5618                        loop {
5619                            let room = max_tokens.saturating_sub(generated);
5620                            if room <= 2 {
5621                                break;
5622                            }
5623                            let k = burst_k.min(room - 1);
5624                            if k < 1 {
5625                                break;
5626                            }
5627                            let Some(ids) = self.try_multi_burst(t_fwd, next_pos, k) else {
5628                                if self
5629                                    .graph_failed
5630                                    .swap(false, std::sync::atomic::Ordering::Relaxed)
5631                                {
5632                                    self.finish_generation(&mut mtp, &mut router, true);
5633                                    return Err(
5634                                        "GPU token graph failed during greedy burst".to_string()
5635                                    );
5636                                }
5637                                break;
5638                            };
5639                            next_pos += k;
5640                            for &id in &ids {
5641                                if !commit!(id) {
5642                                    stopped = true;
5643                                    break;
5644                                }
5645                            }
5646                            if stopped {
5647                                break;
5648                            }
5649                            t_fwd = *ids.last().unwrap();
5650                        }
5651                        if stopped {
5652                            break 'decode;
5653                        }
5654                    }
5655                    // Metal: keep the draft head's cache in step through
5656                    // the trial's plain phase and a paused speculation —
5657                    // the pair (hidden, t_fwd) at next_pos−1, the step the
5658                    // round's draft 0 would take. Without it the head's
5659                    // cache lagged the trunk by every plain token for the
5660                    // rest of the generation: the batched warm-up declined
5661                    // every later round and its rows went one by one (a
5662                    // whole MTP step per accepted token), and the drafts
5663                    // attended a context with those tokens missing.
5664                    #[cfg(target_os = "macos")]
5665                    if graph_spec
5666                        && spec_watchdog_off
5667                        && next_pos > 0
5668                        && self.mtp_graph_mode == Some(true)
5669                        && crate::gpu::q1_force()
5670                    {
5671                        if let Some(m) = mtp.as_mut() {
5672                            let _ = self.mtp_step_metal(m, &hidden, t_fwd, next_pos - 1, false);
5673                        }
5674                    }
5675                    hidden = self.forward_layers(&self.embed_single(t_fwd), next_pos, task_mask);
5676                    next_pos += 1;
5677                    // Dynamic routing: the forward updated φ; ask the
5678                    // router whether to switch skills before the next token.
5679                    if let Some(r) = &mut router {
5680                        let phi = self.dyn_phi_ema.clone();
5681                        let decision = r.step(&phi, generated);
5682                        if let Some(new_active) = decision {
5683                            let _ = self.set_active_skill(new_active);
5684                        }
5685                        // Backfill this token's coherence + switch flag from
5686                        // the just-run eval (freshest measured values).
5687                        if trace_on {
5688                            if let Some(last) = traces.last_mut() {
5689                                let e = r.last_best_e();
5690                                last.recon = e.is_finite().then_some(e);
5691                                last.switched = decision.is_some();
5692                            }
5693                        }
5694                    }
5695                }
5696            }
5697        }
5698
5699        let cancelled = finish_reason == "cancelled";
5700        // A generation during which the router switched weights holds no
5701        // state any single overlay would produce (each switch cleared the
5702        // cache mid-sequence), so it leaves no reuse key behind either.
5703        let dyn_switched = router.as_ref().is_some_and(|r| !r.switches.is_empty());
5704        if mimo_spec {
5705            if let Some(st) = self.mimo_mtp.as_ref() {
5706                let line = st.stats.line();
5707                tracing::info!("{line}");
5708                if std::env::var_os("CMF_MIMO_MTP_STATS").is_some() {
5709                    eprintln!("{line}");
5710                }
5711            }
5712        }
5713        self.finish_generation(&mut mtp, &mut router, cancelled);
5714
5715        let output_ids = &all_ids[input_ids.len()..];
5716        // Forwarded = prompt + all generated but the LAST sampled token
5717        // (emitted without being fed back). Exact only without MTP —
5718        // reuse is gated off when MTP is active.
5719        let forwarded = input_ids.len() + output_ids.len().saturating_sub(1);
5720        // A MiMo speculative round that stopped on an accepted draft (EOS,
5721        // cancel) leaves verify rows past the committed stream in the cache:
5722        // never offer that cache for reuse. Neither does a router that
5723        // switched weights mid-sequence (no single overlay produced it).
5724        let consumed = std::mem::take(&mut all_ids);
5725        if dyn_switched {
5726            self.clear_sequence_state();
5727        } else if cancelled || mimo_spec || prompt_rows.is_some() {
5728            self.clear_history();
5729        } else {
5730            self.record_consumed_prefix(&consumed[..forwarded.min(consumed.len())], reuse_from);
5731        }
5732        all_ids = consumed;
5733        let output_ids = &all_ids[input_ids.len()..];
5734        confidence.truncate(output_ids.len()); // guard against any overshoot
5735        traces.truncate(output_ids.len());
5736        Ok(GenerateResult {
5737            text: self.tokenizer.decode(output_ids),
5738            token_ids: output_ids.to_vec(),
5739            prompt_tokens: input_ids.len(),
5740            tokens_generated: generated,
5741            finish_reason,
5742            mtp_drafted: drafted,
5743            mtp_accepted: accepted,
5744            token_confidence: confidence,
5745            traces,
5746        })
5747    }
5748
5749    /// One MTP step: feed `(hidden_p, token_{p+1})` into the draft head,
5750    /// advance its KV cache at position `p`, return the drafted token
5751    /// for position `p+2`.
5752    fn mtp_step(
5753        &mut self,
5754        m: &mut MtpModule,
5755        hidden: &[f32],
5756        next_token: u32,
5757        position: usize,
5758    ) -> u32 {
5759        self.mtp_step_h(m, hidden, next_token, position).0
5760    }
5761
5762    /// Tally for `CMF_MTP_CHAIN_PROBE`: per depth, how often the CHAIN is
5763    /// still an exact prefix of the real continuation. Printed every 128
5764    /// depth-0 samples so a killed run still shows its table.
5765    fn chain_probe_note(depth: usize, prefix_ok: bool) {
5766        use std::sync::Mutex;
5767        static T: Mutex<Vec<(u64, u64)>> = Mutex::new(Vec::new());
5768        let mut t = T.lock().unwrap();
5769        if t.len() <= depth {
5770            t.resize(depth + 1, (0, 0));
5771        }
5772        t[depth].0 += 1;
5773        t[depth].1 += prefix_ok as u64;
5774        if depth == 0 && t[0].0 % 128 == 0 {
5775            let line: Vec<String> = t
5776                .iter()
5777                .enumerate()
5778                .map(|(d, (n, k))| {
5779                    format!(
5780                        "d{}={:.0}%({n})",
5781                        d + 1,
5782                        100.0 * *k as f64 / (*n).max(1) as f64
5783                    )
5784                })
5785                .collect();
5786            eprintln!("mtp-chain: {}", line.join(" "));
5787        }
5788    }
5789
5790    /// `mtp_step` that also hands back the block's own output hidden — the
5791    /// state a CHAINED draft feeds the next step, the way a multi-token
5792    /// speculative round iterates the head on itself.
5793    /// One MTP block step from (trunk hidden, token): the head's LOGITS
5794    /// and the block's own hidden for chaining. The draft is argmax of the
5795    /// logits on the greedy path and a draw from their post-chain
5796    /// distribution on the sampling path.
5797    fn mtp_step_hl(
5798        &mut self,
5799        m: &mut MtpModule,
5800        hidden: &[f32],
5801        next_token: u32,
5802        position: usize,
5803    ) -> (Vec<f32>, Vec<f32>) {
5804        // The graph arm: the MTP block as a one-layer token graph with the
5805        // head fused — device attention over the block's own KV mirror,
5806        // one submit for block + head, hidden and logits back together.
5807        // Decided once per generation (see `mtp_graph_mode`).
5808        #[cfg(target_os = "macos")]
5809        if self.mtp_graph_mode != Some(false) && crate::gpu::q1_force() {
5810            if let Some(r) = self.mtp_step_metal(m, hidden, next_token, position, true) {
5811                self.mtp_graph_mode = Some(true);
5812                return r;
5813            }
5814            if self.mtp_graph_mode == Some(true) {
5815                tracing::error!("mtp Metal graph failed after admission");
5816                self.clear_sequence_state();
5817                self.graph_failed
5818                    .store(true, std::sync::atomic::Ordering::Relaxed);
5819                self.cancel
5820                    .store(true, std::sync::atomic::Ordering::Relaxed);
5821                return (Vec::new(), Vec::new());
5822            }
5823            self.mtp_graph_mode = Some(false);
5824        }
5825        #[cfg(feature = "gpu")]
5826        if self.mtp_graph_mode != Some(false) {
5827            if !self.mtp_graph_ok(m) {
5828                if self.mtp_graph_mode == Some(true) {
5829                    // A mirror was already admitted, so a capability change
5830                    // cannot safely switch this request to the stale CPU
5831                    // cache.  Keep the same terminal contract as a failed
5832                    // token graph.
5833                    tracing::error!("mtp graph became unavailable after admission");
5834                    self.clear_sequence_state();
5835                    self.graph_failed
5836                        .store(true, std::sync::atomic::Ordering::Relaxed);
5837                    self.cancel
5838                        .store(true, std::sync::atomic::Ordering::Relaxed);
5839                    return (Vec::new(), Vec::new());
5840                }
5841                self.mtp_graph_mode = Some(false);
5842            } else {
5843                if let Some(r) = self.mtp_step_graph(m, hidden, next_token, position) {
5844                    self.mtp_graph_mode = Some(true);
5845                    return r;
5846                }
5847                if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
5848                    // A token graph can have admitted a persistent MTP/GDN
5849                    // mirror before its readback failed.  The CPU MTP cache
5850                    // is not a valid continuation in that state; leave the
5851                    // flag set so the generation caller returns through its
5852                    // terminal error path instead of silently switching
5853                    // arithmetic.
5854                    return (Vec::new(), Vec::new());
5855                }
5856                // `mtp_graph_ok` was true, so a None here means a refusal or
5857                // failure after graph admission.  Do not fall through to a
5858                // CPU cache whose rows may lag the device mirror.
5859                tracing::error!("mtp graph failed or declined after admission");
5860                self.clear_sequence_state();
5861                self.graph_failed
5862                    .store(true, std::sync::atomic::Ordering::Relaxed);
5863                self.cancel
5864                    .store(true, std::sync::atomic::Ordering::Relaxed);
5865                return (Vec::new(), Vec::new());
5866            }
5867        }
5868        // fc concat order is [enorm(embed); hnorm(hidden)] — EMBEDDING
5869        // FIRST. Verified by the oracle (converter/mtp_oracle.py):
5870        // [emb;hid] → 45.8% acceptance, [hid;emb] → 0.00%.
5871        let e = self.embed_single(next_token);
5872        let mut cat = vec![0.0f32; 2 * self.hidden_size];
5873        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
5874        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
5875        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
5876        let mut x = vec![0.0f32; self.hidden_size];
5877        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
5878
5879        // One standard transformer block over the MTP's own cache.
5880        let lw = &m.layer;
5881        inference::rms_norm_into(
5882            &x,
5883            &lw.input_norm,
5884            self.rms_eps,
5885            self.norm_style,
5886            &mut self.ws.n1,
5887        );
5888        let attn = match &lw.attn {
5889            // MLA models carry no MTP head; this path cannot see them.
5890            AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
5891            AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
5892            AttnKind::Full {
5893                wq,
5894                wk,
5895                wv,
5896                wo,
5897                q_norm,
5898                k_norm,
5899                output_gate,
5900                softplus_gate,
5901                bias,
5902            } => {
5903                let mut cfg = self.attn_cfg(position);
5904                cfg.q_norm = q_norm.as_deref();
5905                cfg.k_norm = k_norm.as_deref();
5906                cfg.output_gate = *output_gate;
5907                cfg.softplus_gate = softplus_gate
5908                    .as_ref()
5909                    .map(|(gate, per_head)| (gate, *per_head));
5910                cfg.bias = bias
5911                    .as_ref()
5912                    .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
5913                attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
5914            }
5915            AttnKind::Linear(_)
5916            | AttnKind::LinearGdn(_)
5917            | AttnKind::ShortConv(_)
5918            | AttnKind::Bounded(_) => {
5919                unreachable!("MTP block is full attention")
5920            }
5921        };
5922        for (i, &a) in attn.iter().enumerate() {
5923            x[i] += a;
5924        }
5925        inference::rms_norm_into(
5926            &x,
5927            &lw.post_norm,
5928            self.rms_eps,
5929            self.norm_style,
5930            &mut self.ws.p1,
5931        );
5932        let ffn = ffn_forward(&lw.ffn, &self.ws.p1, self.pool.as_deref(), None);
5933        for (i, &f) in ffn.iter().enumerate() {
5934            x[i] += f;
5935        }
5936
5937        inference::rms_norm_into(
5938            &x,
5939            &m.final_norm,
5940            self.rms_eps,
5941            self.norm_style,
5942            &mut self.ws.n1,
5943        );
5944        let lg = self.lm_head_forward(&self.ws.n1);
5945        (lg, x)
5946    }
5947
5948    /// `mtp_step_hl` reduced to the greedy draft: argmax of the head.
5949    fn mtp_step_h(
5950        &mut self,
5951        m: &mut MtpModule,
5952        hidden: &[f32],
5953        next_token: u32,
5954        position: usize,
5955    ) -> (u32, Vec<f32>) {
5956        let (mut lg, x) = self.mtp_step_hl(m, hidden, next_token, position);
5957        let draft = sampler::argmax(&lg);
5958        attention::recycle_buf(&mut lg);
5959        (draft, x)
5960    }
5961
5962    /// One speculative round for the trial: rounds 1..5 of a `Spec` phase
5963    /// advance it (the monitor already averaged this round); after five,
5964    /// the plain phase runs (once — a known plain rate decides at once);
5965    /// a decided speculation keeps re-checking the rule every round and
5966    /// stops after four losing rounds in a row.
5967    fn spec_trial_round(trial: SpecTrial, mon: &mut SpecMon, generated: usize) -> SpecTrial {
5968        match trial {
5969            SpecTrial::Spec { t0, gen0, rounds } => {
5970                let rounds = rounds + 1;
5971                if rounds >= 5 {
5972                    if mon.plain_ms > 0.0 {
5973                        let keep = mon.pays();
5974                        mon.fails = 0;
5975                        tracing::info!(
5976                            "speculation re-check: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok — {}",
5977                            mon.tokens,
5978                            mon.round_ms,
5979                            mon.plain_ms,
5980                            if keep { "speculating" } else { "plain" }
5981                        );
5982                        SpecTrial::Decided {
5983                            spec: keep,
5984                            recheck_at: if keep { usize::MAX } else { generated + 128 },
5985                        }
5986                    } else if mon.pays() {
5987                        // Metal: the rounds land enough tokens each that no
5988                        // plain measurement is needed — keep speculating,
5989                        // and re-check every round (a losing streak sends
5990                        // the loop to the plain phase, below).
5991                        mon.fails = 0;
5992                        tracing::info!(
5993                            "speculation trial: {:.2} tok/round in {:.1} ms — speculating (plain not timed)",
5994                            mon.tokens,
5995                            mon.round_ms,
5996                        );
5997                        SpecTrial::Decided {
5998                            spec: true,
5999                            recheck_at: usize::MAX,
6000                        }
6001                    } else {
6002                        SpecTrial::Plain {
6003                            t0: std::time::Instant::now(),
6004                            gen0: generated,
6005                        }
6006                    }
6007                } else {
6008                    SpecTrial::Spec { t0, gen0, rounds }
6009                }
6010            }
6011            SpecTrial::Decided { spec: true, .. } => {
6012                if mon.pays() {
6013                    mon.fails = 0;
6014                    trial
6015                } else {
6016                    mon.fails += 1;
6017                    if mon.fails >= 4 {
6018                        if mon.plain_ms <= 0.0 {
6019                            // Metal, plain never timed: four doubtful rounds
6020                            // buy the (bounded) plain measurement, and the
6021                            // exact rule decides from it.
6022                            tracing::info!(
6023                                "speculation doubtful: {:.2} tok/round in {:.1} ms — timing plain",
6024                                mon.tokens,
6025                                mon.round_ms,
6026                            );
6027                            return SpecTrial::Plain {
6028                                t0: std::time::Instant::now(),
6029                                gen0: generated,
6030                            };
6031                        }
6032                        tracing::info!(
6033                            "speculation stopped: {:.2} tok/round in {:.1} ms vs plain {:.1} ms/tok",
6034                            mon.tokens,
6035                            mon.round_ms,
6036                            mon.plain_ms
6037                        );
6038                        SpecTrial::Decided {
6039                            spec: false,
6040                            recheck_at: generated + 128,
6041                        }
6042                    } else {
6043                        trial
6044                    }
6045                }
6046            }
6047            other => other,
6048        }
6049    }
6050
6051    /// The MTP block's device-mirror id: the trunk's id with a high bit,
6052    /// so the (kv_id, layer) mirror keys never collide.
6053    fn mtp_kv_id(&self) -> u64 {
6054        self.graph_kv_id | (1u64 << 40)
6055    }
6056
6057    /// The MTP block's mirror layer index: 0 — its own kv_id keeps it
6058    /// apart from the trunk, and the BATCH graph (the warm-up path) keys
6059    /// its mirrors at layer 0 with no base of its own, so the draft's
6060    /// token graph must key the same slot.
6061    const MTP_LAYER_BASE: usize = 0;
6062
6063    /// The wgpu MTP draft writes speculative rows straight into its device
6064    /// mirror while the CPU owner retains only the real prompt/decode anchor.
6065    /// After verification, move that mirror cursor back to the anchor before
6066    /// replaying accepted pairs.  The next graph append then sees the same
6067    /// contiguous position as the CPU/Metal path without uploading stale
6068    /// speculative rows.
6069    #[cfg(feature = "gpu")]
6070    fn rewind_mtp_graph_mirror(&self, stored: usize) -> bool {
6071        self.mtp_graph_mode != Some(true)
6072            || crate::gpu::graph_kv_set_stored(self.mtp_kv_id(), Self::MTP_LAYER_BASE, stored)
6073    }
6074
6075    /// A speculative verify graph appends the full `k+1` trunk rows before
6076    /// the acceptance count is known.  GDN state already has a snapshot
6077    /// restore; Full-attention mirrors need the matching logical cursor
6078    /// rewind so the next graph call does not reject an ahead-of-position KV
6079    /// cache after a partial acceptance.
6080    #[cfg(feature = "gpu")]
6081    fn rewind_trunk_graph_mirrors(&self, stored: usize) -> bool {
6082        let mut ok = true;
6083        let mut expected = false;
6084        for li in 0..self.num_layers {
6085            if matches!(
6086                self.weights.layers[self.phys_layer(li)].attn,
6087                AttnKind::Full { .. }
6088            ) {
6089                expected = true;
6090                ok &= crate::gpu::graph_kv_set_stored(self.graph_kv_id, li, stored);
6091            }
6092        }
6093        !expected || ok
6094    }
6095
6096    /// Count the recurrent layers participating in the trunk verify graph.
6097    /// Snapshot restore is all-or-nothing across that set; deriving the count
6098    /// from the model keeps the restore contract valid for looped models too.
6099    fn graph_gdn_layer_count(&self) -> usize {
6100        (0..self.num_layers)
6101            .filter(|&li| {
6102                matches!(
6103                    &self.weights.layers[self.phys_layer(li)].attn,
6104                    AttnKind::LinearGdn(_)
6105                )
6106            })
6107            .count()
6108    }
6109
6110    /// The block's input from (trunk hidden, token): eh_proj · [enorm(e);
6111    /// hnorm(h)] — the same arithmetic the per-op path starts with.
6112    fn mtp_block_input(&mut self, m: &MtpModule, hidden: &[f32], next_token: u32) -> Vec<f32> {
6113        let e = self.embed_single(next_token);
6114        let mut cat = vec![0.0f32; 2 * self.hidden_size];
6115        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6116        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6117        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6118        let mut x = vec![0.0f32; self.hidden_size];
6119        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6120        x
6121    }
6122
6123    /// Is the MTP block graphable at all (device up, full attention
6124    /// without softplus, dense FFN)? The plan itself is built per call.
6125    #[cfg(feature = "gpu")]
6126    fn mtp_block_graph_ok(&self, m: &MtpModule) -> bool {
6127        if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0") {
6128            return false;
6129        }
6130        if !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
6131            || !crate::gpu::enabled_here()
6132            || self.attn_softcap > 0.0
6133            || self.attention_heads_per_layer.is_some()
6134            // The block graph caches V as wide as K and feeds o_proj
6135            // nh·head_dim; a narrow-V model keeps its MTP block per-op.
6136            || self.v_head_dim.is_some()
6137        {
6138            return false;
6139        }
6140        matches!(
6141            &m.layer.attn,
6142            AttnKind::Full {
6143                softplus_gate: None,
6144                ..
6145            }
6146        ) && matches!(&m.layer.ffn, FfnKind::Dense(_))
6147    }
6148
6149    /// Full MTP token-graph eligibility, including the fused lm-head and all
6150    /// block projection weights.  Keep this distinct from the block-only
6151    /// check: prompt warm-up does not need the head, while a draft step does.
6152    #[cfg(feature = "gpu")]
6153    fn mtp_graph_ok(&self, m: &MtpModule) -> bool {
6154        if !self.mtp_block_graph_ok(m) {
6155            return false;
6156        }
6157        let AttnKind::Full { wq, wk, wv, wo, .. } = &m.layer.attn else {
6158            return false;
6159        };
6160        let FfnKind::Dense(d) = &m.layer.ffn else {
6161            return false;
6162        };
6163        d.segs.is_empty()
6164            && wq.graph_weight().is_some()
6165            && wk.graph_weight().is_some()
6166            && wv.graph_weight().is_some()
6167            && wo.graph_weight().is_some()
6168            && d.gate_proj.graph_weight().is_some()
6169            && d.up_proj.graph_weight().is_some()
6170            && d.down_proj.graph_weight().is_some()
6171            && self.weights.lm_head.graph_weight().is_some()
6172    }
6173
6174    /// One MTP block step on the wgpu token graph: block + fused head in
6175    /// one submit, the block hidden and the logits read back together.
6176    /// None = the graph cannot take this block (softplus gate, non-dense
6177    /// FFN, unquantized head, no device) — the caller keeps the per-op
6178    /// path for the whole generation.
6179    #[cfg(feature = "gpu")]
6180    fn mtp_step_graph(
6181        &mut self,
6182        m: &mut MtpModule,
6183        hidden: &[f32],
6184        next_token: u32,
6185        position: usize,
6186    ) -> Option<(Vec<f32>, Vec<f32>)> {
6187        if !self.mtp_graph_ok(m) {
6188            return None;
6189        }
6190        let lw = &m.layer;
6191        let AttnKind::Full {
6192            wq,
6193            wk,
6194            wv,
6195            wo,
6196            q_norm,
6197            k_norm,
6198            output_gate,
6199            softplus_gate,
6200            bias,
6201        } = &lw.attn
6202        else {
6203            return None;
6204        };
6205        if softplus_gate.is_some() {
6206            return None;
6207        }
6208        let FfnKind::Dense(d) = &lw.ffn else {
6209            return None;
6210        };
6211        if !d.segs.is_empty() {
6212            return None; // tube layers run on the segmented path
6213        }
6214        // The block's input first: it borrows `self` mutably (embed scratch,
6215        // pool), the plan below borrows the weights immutably.
6216        let mut x = self.mtp_block_input(m, hidden, next_token);
6217        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6218            let (_, i, kind, rs) = t.graph_weight()?;
6219            Some(crate::gpu::GraphW {
6220                idx: i,
6221                kind,
6222                row_scale: rs,
6223                data: &[],
6224                prism: crate::gpu::GraphPrismOp::None,
6225                affine: false,
6226            })
6227        }
6228        let (model, _, _, _) = wq.graph_weight()?;
6229        let model = model.clone();
6230        let (lm_gw, lm_rows) = {
6231            let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6232            // The draft's head over the CMF_DRAFT_VOCAB shortlist (the same
6233            // cut the native Metal draft takes): 662 MB a step on Qwen3.8
6234            // becomes 170 MB at 65536; the verify keeps the full head.
6235            let rows = if kind == 6 {
6236                self.draft_head_rows(self.weights.lm_head.rows())
6237            } else {
6238                self.weights.lm_head.rows()
6239            };
6240            (
6241                crate::gpu::GraphW {
6242                    idx: i,
6243                    kind,
6244                    row_scale: rs,
6245                    data: &[],
6246                    prism: crate::gpu::GraphPrismOp::None,
6247                    affine: false,
6248                },
6249                rows,
6250            )
6251        };
6252        let layer = crate::gpu::GraphLayer {
6253            input_norm: &lw.input_norm,
6254            attn: crate::gpu::GraphAttn::Full {
6255                wq: gw(wq)?,
6256                wk: gw(wk)?,
6257                wv: gw(wv)?,
6258                wo: gw(wo)?,
6259                q_norm: q_norm.as_deref(),
6260                k_norm: k_norm.as_deref(),
6261                late_qk_norm: self.qk_norm_after_rope,
6262                bias: bias
6263                    .as_ref()
6264                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6265                output_gate: *output_gate,
6266                cpu_k: m.kv.k_heads(),
6267                cpu_v: m.kv.v_heads(),
6268                cpu_base: m.kv.base(),
6269                geom: None,
6270                head_gate: None,
6271            },
6272            post_norm: &lw.post_norm,
6273            ffn: crate::gpu::GraphFfn::Dense {
6274                gate: gw(&d.gate_proj)?,
6275                up: gw(&d.up_proj)?,
6276                down: gw(&d.down_proj)?,
6277                act: d.act.graph_act()?,
6278            },
6279        };
6280        let nh = self.num_heads;
6281        let (nkv, hd, rd) = self.layer_geom(0);
6282        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6283        let mut logits = Vec::new();
6284        let ok = crate::gpu::forward_token_graph(
6285            &model,
6286            self.mtp_kv_id(),
6287            std::slice::from_ref(&layer),
6288            &[None],
6289            self.o1_epoch,
6290            &self.inv_freq,
6291            &mut x,
6292            nh,
6293            nkv,
6294            hd,
6295            self.attn_scale,
6296            rd,
6297            self.hidden_size,
6298            self.intermediate_size,
6299            position,
6300            self.kv_cache.max_seq_len,
6301            gemma,
6302            self.rms_eps as f32,
6303            Some((&lm_gw, lm_rows)),
6304            &m.final_norm,
6305            &mut logits,
6306            &[],
6307            1,
6308            None,
6309            None,
6310            None,
6311            Self::MTP_LAYER_BASE,
6312            true,
6313        );
6314        match ok {
6315            crate::gpu::TokenGraphOutcome::Completed => {}
6316            crate::gpu::TokenGraphOutcome::Declined => return None,
6317            crate::gpu::TokenGraphOutcome::Failed => {
6318                // The backend has already admitted persistent state.  Keep
6319                // this distinct from a capability refusal so the caller
6320                // cannot switch to the stale CPU MTP cache.
6321                self.clear_sequence_state();
6322                self.graph_failed
6323                    .store(true, std::sync::atomic::Ordering::Relaxed);
6324                self.cancel
6325                    .store(true, std::sync::atomic::Ordering::Relaxed);
6326                return None;
6327            }
6328        }
6329        logits.resize(self.vocab_size, 0.0);
6330        Some((logits, x))
6331    }
6332
6333    /// The warm-ups of one speculative round on the device: every accepted
6334    /// (hidden, token) pair as ONE batched graph run over the MTP block
6335    /// (no head) — its kv_append lands the pairs in the block's mirror.
6336    /// `pairs` are consecutive positions from `first_pos`.  The tri-state
6337    /// result is intentional: a refusal before admission may use the
6338    /// per-row/CPU route, while a failure after admission must terminate the
6339    /// sequence rather than fall through to a stale CPU cache.
6340    #[cfg(feature = "gpu")]
6341    fn mtp_warm_graph(
6342        &mut self,
6343        m: &mut MtpModule,
6344        pairs: &[(&[f32], u32)],
6345        first_pos: usize,
6346    ) -> crate::gpu::BatchGraphOutcome {
6347        if pairs.is_empty() {
6348            return crate::gpu::BatchGraphOutcome::Completed;
6349        }
6350        if !self.mtp_block_graph_ok(m) {
6351            return crate::gpu::BatchGraphOutcome::Declined;
6352        }
6353        let hs = self.hidden_size;
6354        // Block inputs for every pair (eh_proj on the per-op path, one
6355        // matvec each — the plan's own prologue).
6356        let mut hiddens = Vec::with_capacity(pairs.len() * hs);
6357        for (h, t) in pairs {
6358            hiddens.extend_from_slice(&self.mtp_block_input(m, h, *t));
6359        }
6360        let lw = &m.layer;
6361        let AttnKind::Full {
6362            wq,
6363            wk,
6364            wv,
6365            wo,
6366            q_norm,
6367            k_norm,
6368            output_gate,
6369            bias,
6370            ..
6371        } = &lw.attn
6372        else {
6373            return crate::gpu::BatchGraphOutcome::Declined;
6374        };
6375        let FfnKind::Dense(d) = &lw.ffn else {
6376            return crate::gpu::BatchGraphOutcome::Declined;
6377        };
6378        if !d.segs.is_empty() {
6379            return crate::gpu::BatchGraphOutcome::Declined; // tube layers run on the segmented path
6380        }
6381        // The graph has no arm for a projected output gate, and only the
6382        // activations `graph_act` names.
6383        let (AttnKind::Full { softplus_gate: None, .. }, Some(gact)) =
6384            (&lw.attn, d.act.graph_act())
6385        else {
6386            return crate::gpu::BatchGraphOutcome::Declined;
6387        };
6388        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
6389            let (_, i, kind, rs) = t.graph_weight()?;
6390            Some(crate::gpu::GraphW {
6391                idx: i,
6392                kind,
6393                row_scale: rs,
6394                data: &[],
6395                prism: crate::gpu::GraphPrismOp::None,
6396                affine: false,
6397            })
6398        }
6399        let Some((model, _, _, _)) = wq.graph_weight() else {
6400            return crate::gpu::BatchGraphOutcome::Declined;
6401        };
6402        let model = model.clone();
6403        let (Some(gwq), Some(gwk), Some(gwv), Some(gwo), Some(gg), Some(gu), Some(gd)) = (
6404            gw(wq),
6405            gw(wk),
6406            gw(wv),
6407            gw(wo),
6408            gw(&d.gate_proj),
6409            gw(&d.up_proj),
6410            gw(&d.down_proj),
6411        ) else {
6412            return crate::gpu::BatchGraphOutcome::Declined;
6413        };
6414        let layer = crate::gpu::GraphLayer {
6415            input_norm: &lw.input_norm,
6416            attn: crate::gpu::GraphAttn::Full {
6417                wq: gwq,
6418                wk: gwk,
6419                wv: gwv,
6420                wo: gwo,
6421                q_norm: q_norm.as_deref(),
6422                k_norm: k_norm.as_deref(),
6423                late_qk_norm: self.qk_norm_after_rope,
6424                bias: bias
6425                    .as_ref()
6426                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
6427                output_gate: *output_gate,
6428                cpu_k: m.kv.k_heads(),
6429                cpu_v: m.kv.v_heads(),
6430                cpu_base: m.kv.base(),
6431                geom: None,
6432                head_gate: None,
6433            },
6434            post_norm: &lw.post_norm,
6435            ffn: crate::gpu::GraphFfn::Dense {
6436                gate: gg,
6437                up: gu,
6438                down: gd,
6439                act: gact,
6440            },
6441        };
6442        let positions: Vec<usize> = (first_pos..first_pos + pairs.len()).collect();
6443        let nh = self.num_heads;
6444        let (nkv, hd, rd) = self.layer_geom(0);
6445        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
6446        crate::gpu::forward_batch_graph(
6447            &model,
6448            self.mtp_kv_id(),
6449            std::slice::from_ref(&layer),
6450            &self.inv_freq,
6451            &mut hiddens,
6452            nh,
6453            nkv,
6454            hd,
6455            rd,
6456            hs,
6457            self.intermediate_size,
6458            &positions,
6459            self.kv_cache.max_seq_len,
6460            gemma,
6461            self.rms_eps as f32,
6462            self.attn_scale,
6463            pairs.len(),
6464            &[],
6465            0,
6466            None,
6467            None,
6468        )
6469    }
6470
6471    /// Complete an MTP warm-up after the batched graph has refused.  A
6472    /// graphable block is retried one row at a time; once any device row has
6473    /// been admitted, a CPU fallback would observe a stale mirror, so every
6474    /// token-graph refusal is terminal.  If the block is not graphable and no
6475    /// mirror exists yet, warming on the CPU is safe and records the CPU mode
6476    /// for the rest of the generation.
6477    #[cfg(feature = "gpu")]
6478    fn mtp_warm_graph_fallback(
6479        &mut self,
6480        m: &mut MtpModule,
6481        pairs: &[(&[f32], u32)],
6482        first_pos: usize,
6483    ) -> bool {
6484        if pairs.is_empty() {
6485            return true;
6486        }
6487        let graphable = self.mtp_block_graph_ok(m);
6488        if !graphable {
6489            // A previously admitted mirror cannot be made coherent by
6490            // appending to the host cache.  The caller turns this into a
6491            // terminal generation error and clears both mirrors.
6492            if self.mtp_graph_mode == Some(true) {
6493                return false;
6494            }
6495            self.mtp_graph_mode = Some(false);
6496            for (j, (h, t)) in pairs.iter().enumerate() {
6497                self.mtp_warm(m, h, *t, first_pos + j);
6498            }
6499            return true;
6500        }
6501
6502        // The batch refusal is recoverable only through the same device
6503        // state.  Keep rows owned until each token graph has completed; a
6504        // None is treated as unsafe because the token-graph API deliberately
6505        // collapses its backend refusal/failure into that result.
6506        for (j, (h, t)) in pairs.iter().enumerate() {
6507            if self.mtp_step_graph(m, h, *t, first_pos + j).is_none() {
6508                return false;
6509            }
6510        }
6511        self.mtp_graph_mode = Some(true);
6512        true
6513    }
6514
6515    /// Warm a contiguous set of MTP pairs using the existing graph seam, with
6516    /// an all-or-nothing error contract for callers that already admitted the
6517    /// trunk batch.  The non-GPU build keeps the same pair accounting while
6518    /// using the established CPU warm path.
6519    #[cfg(feature = "gpu")]
6520    fn mtp_warm_prefill_pairs(
6521        &mut self,
6522        m: &mut MtpModule,
6523        pairs: &[(&[f32], u32)],
6524        first_pos: usize,
6525    ) -> Result<(), &'static str> {
6526        // Keep unsupported token-graph heads on the established CPU MTP
6527        // route before admitting any block mirror.  Once a device mirror is
6528        // active, the same condition is terminal because CPU rows cannot
6529        // repair its state.
6530        if self.mtp_graph_mode == Some(false) || !self.mtp_graph_ok(m) {
6531            if self.mtp_graph_mode == Some(true) {
6532                return Err("MTP token graph became unavailable after admission");
6533            }
6534            self.mtp_graph_mode = Some(false);
6535            for (j, (h, t)) in pairs.iter().enumerate() {
6536                self.mtp_warm(m, h, *t, first_pos + j);
6537            }
6538            return Ok(());
6539        }
6540        match self.mtp_warm_graph(m, pairs, first_pos) {
6541            crate::gpu::BatchGraphOutcome::Completed => {
6542                if !pairs.is_empty() {
6543                    self.mtp_graph_mode = Some(true);
6544                }
6545                Ok(())
6546            }
6547            crate::gpu::BatchGraphOutcome::Declined => {
6548                if self.mtp_warm_graph_fallback(m, pairs, first_pos) {
6549                    Ok(())
6550                } else {
6551                    Err("MTP warm-up fallback failed after device admission")
6552                }
6553            }
6554            crate::gpu::BatchGraphOutcome::Failed => {
6555                Err("MTP warm batch graph failed after admission")
6556            }
6557        }
6558    }
6559
6560    #[cfg(not(feature = "gpu"))]
6561    fn mtp_warm_prefill_pairs(
6562        &mut self,
6563        m: &mut MtpModule,
6564        pairs: &[(&[f32], u32)],
6565        first_pos: usize,
6566    ) -> Result<(), &'static str> {
6567        for (j, (h, t)) in pairs.iter().enumerate() {
6568            self.mtp_warm(m, h, *t, first_pos + j);
6569        }
6570        Ok(())
6571    }
6572
6573    /// The MTP block alone — advance its KV with a (hidden, token) pair the
6574    /// verify just proved, without paying the head. What keeps the draft's
6575    /// attention context warm between speculative rounds.
6576    fn mtp_warm(&mut self, m: &mut MtpModule, hidden: &[f32], next_token: u32, position: usize) {
6577        let e = self.embed_single(next_token);
6578        let mut cat = vec![0.0f32; 2 * self.hidden_size];
6579        let (cat_e, cat_h) = cat.split_at_mut(self.hidden_size);
6580        inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
6581        inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
6582        let mut x = vec![0.0f32; self.hidden_size];
6583        m.eh_proj.matvec(&cat, &mut x, self.pool.as_deref());
6584        inference::rms_norm_into(
6585            &x,
6586            &m.layer.input_norm,
6587            self.rms_eps,
6588            self.norm_style,
6589            &mut self.ws.n1,
6590        );
6591        let attn = match &m.layer.attn {
6592            AttnKind::Full {
6593                wq,
6594                wk,
6595                wv,
6596                wo,
6597                q_norm,
6598                k_norm,
6599                output_gate,
6600                softplus_gate,
6601                bias,
6602            } => {
6603                let mut cfg = self.attn_cfg(position);
6604                cfg.q_norm = q_norm.as_deref();
6605                cfg.k_norm = k_norm.as_deref();
6606                cfg.output_gate = *output_gate;
6607                cfg.softplus_gate = softplus_gate.as_ref().map(|(g, p)| (g, *p));
6608                cfg.bias = bias
6609                    .as_ref()
6610                    .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice()));
6611                attention::qwen_attention(&self.ws.n1, wq, wk, wv, wo, &mut m.kv, &cfg)
6612            }
6613            _ => return,
6614        };
6615        let _ = attn;
6616    }
6617
6618    /// Speculative decode ON the wgpu whole-token graph: draft k with the
6619    /// MTP head, verify all of them plus the tip in ONE batched graph
6620    /// submit whose tail folds the head, commit the accepted prefix and
6621    /// roll the GDN state back to the last real position. Greedy only —
6622    /// output equals the plain graph's token for token, the way the DSV4
6623    /// verify equals the walk.
6624    #[cfg(feature = "gpu")]
6625    #[allow(clippy::too_many_arguments)]
6626    fn graph_spec_step(
6627        &mut self,
6628        m: &mut MtpModule,
6629        hidden: &[f32],
6630        t_next: u32,
6631        next_pos: usize,
6632        drafted: &mut usize,
6633        accepted: &mut usize,
6634        // The committed stream (prompt + generated so far, `t_next`
6635        // included): the sampler chain's penalties read it, and the
6636        // sampling arm extends it with the drafts position by position.
6637        all_ids: &mut Vec<u32>,
6638        // Tokens left before `max_tokens`. A round commits up to k
6639        // accepted drafts, and those positions are already in the cache,
6640        // so the depth is capped here — trimming the output afterwards
6641        // would leave cache rows the committed stream does not have.
6642        room: usize,
6643    ) -> Option<(Vec<u32>, usize, Vec<f32>)> {
6644        // 3 is the measured optimum on Qwen3.6-27B / RTX 5090 (medians
6645        // of three, greedy): 51.1 tok/s against a plain 49.4, where k=2
6646        // gives 46.1, k=4 50.0, k=5 47.4, k=6 45.2. Acceptance is 89-91%
6647        // throughout — what turns the curve over is the verify, which
6648        // costs ~7.4 ms per extra position, and the draft ~3 ms a step.
6649        // 4 since the draft moved onto the graph (Qwen3.8-27B / 5090:
6650        // k=3 51.2, k=4 51.8 with the per-op draft; the graph draft
6651        // halves the draft cost, so the extra draft is cheaper still).
6652        // 5 with the int8 verify (the default: measured 76.5 against
6653        // k=4's 72-74 and k=6's 74 on the 5090), 4 with the f32 one.
6654        #[cfg(target_os = "macos")]
6655        let metal_native = crate::gpu::q1_force();
6656        #[cfg(not(target_os = "macos"))]
6657        let metal_native = false;
6658        #[cfg(feature = "gpu")]
6659        let k_default = if metal_native {
6660            // the Metal verify's GEMM tile is 8 rows wide and flat in b:
6661            // seven drafts + the tip fill it for free
6662            7
6663        } else if crate::gpu_wgpu::verify_i8_on() {
6664            5
6665        } else {
6666            4
6667        };
6668        #[cfg(not(feature = "gpu"))]
6669        let k_default = 4;
6670        let k_env: Option<usize> = std::env::var("CMF_GRAPH_SPEC_K")
6671            .ok()
6672            .and_then(|v| v.parse().ok())
6673            .filter(|&v| (1..=8).contains(&v));
6674        // Adaptive depth: start below the card's flat-verify optimum and
6675        // let the accepted fraction move it — predictable text climbs to
6676        // the old default within a few rounds, prose settles at 2-3 where
6677        // the shorter verify pays.
6678        let (k_start, k_max) = if metal_native { (7, 7) } else { (3, k_default.max(5)) };
6679        let k_full: usize = k_env.unwrap_or_else(|| self.spec_k_adapt.unwrap_or(k_start));
6680        let k_spec = k_full.min(room).max(1);
6681        // a tail round cut short by `room` says nothing about the text:
6682        // it must not move the adaptive depth the next request starts at
6683        let k_capped = k_spec < k_full;
6684        if next_pos == 0 {
6685            return None;
6686        }
6687        let t_round = std::time::Instant::now();
6688        // Submissions per phase — and they say where the round's money is.
6689        // Qwen3.6-27B on an RTX 5090, k=3:
6690        //
6691        //   draft   9.3 ms / 12 submissions   (four per MTP step)
6692        //   verify 52.8 ms /  1               (the batched graph)
6693        //   commit  5.4 ms /  6               (two per warm)
6694        //
6695        // The verify is already one submit. The draft's own work is 834 MB
6696        // a step — 0.8 ms at this card's measured 1056 GB/s — against 3.1
6697        // ms measured, so ~0.58 ms of every step is round trip, not
6698        // arithmetic, and the same holds for the warms. Eighteen round
6699        // trips a round at roughly half a millisecond each is ~11 ms of a
6700        // 68 ms round: fusing the MTP block into ONE submit the way the
6701        // trunk already is projects to ~64 tok/s against today's 50.9.
6702        // That is the largest measured item left on this path.
6703        let subs = || crate::gpu_wgpu::SUBMITS.load(std::sync::atomic::Ordering::Relaxed);
6704        let sub0 = subs();
6705        // Greedy without penalties verifies by argmax equality (bit-exact
6706        // against the plain path). Anything else is speculative SAMPLING:
6707        // each draft is a DRAW from the MTP head's post-chain distribution
6708        // q_j, kept for the accept test; the verify's rows give p_j.
6709        let cfg = self.sampler_config.clone();
6710        let penalized = !(cfg.repetition_penalty == 1.0
6711            && cfg.presence_penalty == 0.0
6712            && cfg.suppress_tokens.is_empty());
6713        // Three verify regimes: plain greedy (argmax of the raw rows),
6714        // greedy WITH penalties (argmax of the penalized rows — a single
6715        // pass each, no distributions), and sampling (draw / accept /
6716        // correct on post-chain distributions).
6717        let greedy_pen = cfg.temperature < 1e-6 && penalized;
6718        let sampling = cfg.temperature >= 1e-6;
6719        // Sampling with a top-k goes through the SPARSE chain: the dense
6720        // one builds nine 248k-float distributions a round (four drafts,
6721        // five verify rows) and measured 19-22 tok/s against a plain 40 —
6722        // the host, not the card. Sparse, the same nine cost tens of
6723        // microseconds each.
6724        let sparse = sampling && sampler::sparse_ok(&cfg);
6725        let base_len = all_ids.len();
6726        if sampling && !sparse && self.spec_q.len() < k_spec {
6727            self.spec_q.resize_with(k_spec, Vec::new);
6728        }
6729        if sparse && self.spec_qs.len() < k_spec {
6730            self.spec_qs.resize_with(k_spec, Vec::new);
6731        }
6732        // Draft the chain: first from the trunk's tip hidden, then the head
6733        // iterating on itself. Rows land in the MTP KV; the chain rows past
6734        // the first are speculation over speculative state and roll back
6735        // below, replaced by verified pairs.
6736        let mut drafts = Vec::with_capacity(k_spec);
6737        let mut hx = hidden.to_vec();
6738        // CMF_SPEC_DBG=1: draft 0 through BOTH MTP arms (graph and per-op)
6739        // from the same inputs — are the arms the difference, or the inputs?
6740        let spec_dbg = std::env::var("CMF_SPEC_DBG").is_ok();
6741        spec_stamp("pro");
6742        // Plain greedy on native Metal: the whole chain as one command
6743        // buffer (device argmax + embedding gather between the steps).
6744        // A decline before commit hands the round to the per-step loop
6745        // below; a failure after commit is terminal, like any graph
6746        // failure after admission.
6747        #[cfg(target_os = "macos")]
6748        if metal_native && !sampling && !greedy_pen && self.mtp_graph_mode != Some(false) {
6749            match self.mtp_draft_chain_metal(m, hidden, t_next, next_pos - 1, k_spec) {
6750                Ok(ids) => {
6751                    self.mtp_graph_mode = Some(true);
6752                    drafts = ids;
6753                }
6754                Err(true) => {
6755                    tracing::error!("mtp Metal draft chain failed after commit");
6756                    self.clear_sequence_state();
6757                    self.graph_failed
6758                        .store(true, std::sync::atomic::Ordering::Relaxed);
6759                    self.cancel
6760                        .store(true, std::sync::atomic::Ordering::Relaxed);
6761                    return None;
6762                }
6763                Err(false) => {}
6764            }
6765        }
6766        for j in drafts.len()..k_spec {
6767            let tok_in = if j == 0 { t_next } else { drafts[j - 1] };
6768            let mut dbg_ref: Option<(Vec<f32>, Vec<f32>)> = None;
6769            if spec_dbg {
6770                let saved = self.mtp_graph_mode;
6771                self.mtp_graph_mode = Some(false);
6772                let r = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6773                self.mtp_graph_mode = saved;
6774                if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6775                    return None;
6776                }
6777                m.kv.truncate_last(1);
6778                dbg_ref = Some(r);
6779            }
6780            let (mut lg, hj) = self.mtp_step_hl(m, &hx, tok_in, next_pos - 1 + j);
6781            if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
6782                return None;
6783            }
6784            if let Some((lg_cpu, h_cpu)) = dbg_ref {
6785                let n = |v: &[f32]| v.iter().map(|x| x * x).sum::<f32>().sqrt();
6786                let dl = lg
6787                    .iter()
6788                    .zip(&lg_cpu)
6789                    .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6790                let dh = hj
6791                    .iter()
6792                    .zip(&h_cpu)
6793                    .fold(0f32, |m, (a, b)| m.max((a - b).abs()));
6794                eprintln!(
6795                    "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 {}",
6796                    next_pos - 1 + j,
6797                    sampler::argmax(&lg_cpu),
6798                    sampler::argmax(&lg),
6799                    n(&h_cpu),
6800                    n(&hj),
6801                    m.kv.seq_len
6802                );
6803            }
6804            let dj = if sparse {
6805                let mut q = std::mem::take(&mut self.spec_qs[j]);
6806                let ok = sampler::sparse_distribution_into(
6807                    &lg,
6808                    &cfg,
6809                    all_ids,
6810                    &mut self.sampler_scratch,
6811                    self.pool.as_deref(),
6812                    &mut q,
6813                );
6814                let d = if ok {
6815                    sampler::draw_sparse(&q, &mut self.rng)
6816                } else {
6817                    // everything filtered: the dense chain's greedy fallback
6818                    let t = sampler::argmax(&lg);
6819                    q.clear();
6820                    q.push((t, 1.0));
6821                    t
6822                };
6823                self.spec_qs[j] = q;
6824                all_ids.push(d);
6825                d
6826            } else if sampling {
6827                let mut q = std::mem::take(&mut self.spec_q[j]);
6828                sampler::distribution_into(
6829                    &lg,
6830                    &cfg,
6831                    all_ids,
6832                    &mut self.sampler_scratch,
6833                    self.pool.as_deref(),
6834                    &mut q,
6835                );
6836                let d = sampler::draw(&q, &mut self.rng);
6837                self.spec_q[j] = q;
6838                all_ids.push(d); // the next draft's penalties see this one
6839                d
6840            } else if greedy_pen {
6841                let d = sampler::argmax_penalized(
6842                    &lg,
6843                    &cfg,
6844                    all_ids,
6845                    &mut self.sampler_scratch,
6846                    self.pool.as_deref(),
6847                );
6848                all_ids.push(d);
6849                d
6850            } else {
6851                sampler::argmax(&lg)
6852            };
6853            attention::recycle_buf(&mut lg);
6854            drafts.push(dj);
6855            hx = hj;
6856            spec_stamp("d.pick");
6857        }
6858        all_ids.truncate(base_len);
6859        *drafted += k_spec;
6860        let t_draft = t_round.elapsed();
6861        let sub_draft = subs();
6862        // Verify batch: [t_next, d1 .. d_{k-1}] at next_pos.. — every row's
6863        // logits come back from the graph's own head.
6864        let b = k_spec + 1;
6865        let mut hiddens = vec![0.0f32; b * self.hidden_size];
6866        for (i, &t) in std::iter::once(&t_next).chain(drafts.iter()).enumerate() {
6867            let e = self.embed_single(t);
6868            hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&e);
6869        }
6870        let positions: Vec<usize> = (next_pos..next_pos + b).collect();
6871        spec_stamp("v.emb");
6872        let (lm_gw, lm_rows) = {
6873            let (_, i, kind, rs) = self.weights.lm_head.graph_weight()?;
6874            (
6875                crate::gpu::GraphW {
6876                    idx: i,
6877                    kind,
6878                    row_scale: rs,
6879                    data: &[],
6880                    prism: crate::gpu::GraphPrismOp::None,
6881                    affine: false,
6882                },
6883                self.weights.lm_head.rows(),
6884            )
6885        };
6886        let mut logits = Vec::new();
6887        let final_norm = self.weights.final_norm.clone();
6888        // Plain greedy on Metal: the b argmaxes come from the device
6889        // (`argmax_rows` after the head) and the 7.9 MB logits plane is
6890        // never read back — the round's decision needs only the ids, and
6891        // the loop top takes the last verified id as `spec_forced`, which
6892        // is exactly what its argmax of the row would give. The full rows
6893        // stay for anything that reads them: sampling, penalties,
6894        // confidence, the verify oracle, the logit dump.
6895        // `CMF_METAL_DEV_ARGMAX=0` keeps the host path.
6896        #[cfg(target_os = "macos")]
6897        let greedy_dev = metal_native
6898            && !sampling
6899            && !greedy_pen
6900            && !self.confidence_on
6901            && self.final_softcap.is_none()
6902            // The host acceptance argmax scans the WHOLE head row
6903            // (`lm_rows`), the sampler's own row only `vocab_size`: they
6904            // coincide exactly when the head has no padding rows, and
6905            // only then is the device argmax (which scores `vocab_size`)
6906            // bit-identical to both.
6907            && self.vocab_size == lm_rows
6908            && std::env::var_os("CMF_METAL_VERIFY_CHECK").is_none()
6909            && std::env::var_os("CMF_LOGIT_DUMP").is_none()
6910            && std::env::var("CMF_METAL_DEV_ARGMAX").as_deref() != Ok("0");
6911        #[cfg(not(target_os = "macos"))]
6912        let greedy_dev = false;
6913        let mut dev_ids: Vec<u32> = Vec::new();
6914        #[cfg(target_os = "macos")]
6915        let verify_outcome = if metal_native {
6916            let lm = self.weights.lm_head.q1_parts()?;
6917            let n_score = self.vocab_size.min(lm_rows);
6918            self.try_batch_graph_metal(
6919                &mut hiddens,
6920                &positions,
6921                b,
6922                Some((lm, &final_norm, &mut logits)),
6923                if greedy_dev {
6924                    Some((n_score, &mut dev_ids))
6925                } else {
6926                    None
6927                },
6928            )
6929        } else {
6930            self.try_batch_graph_wgpu(
6931                &mut hiddens,
6932                &positions,
6933                b,
6934                Some(crate::gpu::SpecTail {
6935                    lm: lm_gw,
6936                    lm_rows,
6937                    final_norm: &final_norm,
6938                    logits_out: &mut logits,
6939                }),
6940            )
6941        };
6942        #[cfg(not(target_os = "macos"))]
6943        let verify_outcome = self.try_batch_graph_wgpu(
6944            &mut hiddens,
6945            &positions,
6946            b,
6947            Some(crate::gpu::SpecTail {
6948                lm: lm_gw,
6949                lm_rows,
6950                final_norm: &final_norm,
6951                logits_out: &mut logits,
6952            }),
6953        );
6954        match verify_outcome {
6955            crate::gpu::BatchGraphOutcome::Completed => {}
6956            crate::gpu::BatchGraphOutcome::Declined => {
6957                // The verifier refused before admission.  Its draft MTP
6958                // rows are still device-resident, so rewind the separate
6959                // mirror before the caller takes the exact one-token path.
6960                m.kv.truncate_last(k_spec);
6961                if !metal_native && !self.rewind_mtp_graph_mirror(next_pos) {
6962                    self.clear_sequence_state();
6963                    self.graph_failed
6964                        .store(true, std::sync::atomic::Ordering::Relaxed);
6965                    self.cancel
6966                        .store(true, std::sync::atomic::Ordering::Relaxed);
6967                    tracing::error!("MTP graph mirror rewind failed after verify decline");
6968                }
6969                return None;
6970            }
6971            crate::gpu::BatchGraphOutcome::Failed => {
6972                // A failed batch may have advanced trunk/GDN state.  Clear
6973                // both mirrors and preserve the terminal outcome rather than
6974                // falling through to stale CPU state.
6975                self.clear_sequence_state();
6976                self.graph_failed
6977                    .store(true, std::sync::atomic::Ordering::Relaxed);
6978                self.cancel
6979                    .store(true, std::sync::atomic::Ordering::Relaxed);
6980                tracing::error!("MTP verify batch graph failed after admission");
6981                return None;
6982            }
6983        }
6984        // `CMF_METAL_VERIFY_CHECK=1`: run the same b tokens through the
6985        // plain per-token path and compare each row's argmax + logits with
6986        // the verify's — the bring-up oracle for the batched graph. The
6987        // plain forwards mutate the CPU state; it is snapshotted and put
6988        // back, and the K/V mirrors re-pointed, before the round goes on.
6989        #[cfg(target_os = "macos")]
6990        if metal_native && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("1") {
6991            let snap: Vec<Vec<f32>> = self
6992                .kv_cache
6993                .layers
6994                .iter()
6995                .map(|l| l.linear_state.clone())
6996                .collect();
6997            let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
6998            let toks: Vec<u32> = std::iter::once(t_next)
6999                .chain(drafts.iter().copied())
7000                .collect();
7001            let want_save = self.graph_want_logits;
7002            self.graph_want_logits = false;
7003            for (i, &t) in toks.iter().enumerate() {
7004                let hi = self.forward_layers(&self.embed_single(t), next_pos + i, None);
7005                let _ = self.graph_logits.take();
7006                // CMF_SPEC_PLAIN_HIDDEN=1: the next round drafts from the
7007                // plain path's hidden instead of the verify's (an experiment
7008                // on the chain's sensitivity to the half-GEMM noise)
7009                if std::env::var("CMF_SPEC_PLAIN_HIDDEN").as_deref() == Ok("1") {
7010                    hiddens[i * self.hidden_size..(i + 1) * self.hidden_size].copy_from_slice(&hi);
7011                }
7012                let ref_lg = self.logits_from_hidden(&hi);
7013                let row = &logits[i * lm_rows..(i + 1) * lm_rows];
7014                let ra = sampler::argmax(&ref_lg);
7015                let va = sampler::argmax(row);
7016                let mut md = 0f32;
7017                let mut rms = 0f64;
7018                for j in 0..lm_rows.min(ref_lg.len()) {
7019                    let d = (ref_lg[j] - row[j]).abs();
7020                    md = md.max(d);
7021                    rms += (d as f64) * (d as f64);
7022                }
7023                let mut hd = 0f32;
7024                for j in 0..self.hidden_size {
7025                    hd = hd.max((hi[j] - hiddens[i * self.hidden_size + j]).abs());
7026                }
7027                eprintln!(
7028                    "verify-check row {i} tok {t} pos {}: ref argmax {ra} verify argmax {va} {} | max|dlogit| {md:.3} rms {:.4} | max|dhidden| {hd:.4}",
7029                    next_pos + i,
7030                    if ra == va { "OK" } else { "MISMATCH" },
7031                    (rms / lm_rows as f64).sqrt()
7032                );
7033            }
7034            self.graph_want_logits = want_save;
7035            // restore IN PLACE: the pending verify graph wraps these very
7036            // allocations (zero-copy) — replacing the Vec would strand it
7037            for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7038                if l.linear_state.len() == st.len() {
7039                    l.linear_state.copy_from_slice(&st);
7040                } else {
7041                    l.linear_state = st;
7042                }
7043            }
7044            for (li, (l, n0)) in self.kv_cache.layers.iter_mut().zip(attn_lens).enumerate() {
7045                let extra = l.seq_len.saturating_sub(n0);
7046                if extra > 0 {
7047                    l.truncate_last(extra);
7048                    crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, n0);
7049                }
7050            }
7051        }
7052        let t_verify = t_round.elapsed();
7053        let sub_verify = subs();
7054        // Acceptance. Greedy: row i's argmax is the trunk's token after
7055        // input i. Sampling: accept draft i with min(1, p_i/q_i), and on
7056        // the first rejection draw the correction from max(0, p_i − q_i)
7057        // — that token is committed by the loop top as-is (spec_forced).
7058        let mut a = 0usize;
7059        let mut forced: Option<u32> = None;
7060        let ids: Vec<u32> = if sparse {
7061            let mut p = std::mem::take(&mut self.spec_ps);
7062            let mut res = std::mem::take(&mut self.spec_ress);
7063            while a < k_spec {
7064                let ok = sampler::sparse_distribution_into(
7065                    &logits[a * lm_rows..(a + 1) * lm_rows],
7066                    &cfg,
7067                    all_ids,
7068                    &mut self.sampler_scratch,
7069                    self.pool.as_deref(),
7070                    &mut p,
7071                );
7072                if !ok {
7073                    let t = sampler::argmax(&logits[a * lm_rows..(a + 1) * lm_rows]);
7074                    p.clear();
7075                    p.push((t, 1.0));
7076                }
7077                match sampler::spec_accept_or_correct_sparse(
7078                    &p,
7079                    &self.spec_qs[a],
7080                    drafts[a],
7081                    &mut self.rng,
7082                    &mut res,
7083                ) {
7084                    None => {
7085                        all_ids.push(drafts[a]);
7086                        a += 1;
7087                    }
7088                    Some(c) => {
7089                        forced = Some(c);
7090                        break;
7091                    }
7092                }
7093            }
7094            all_ids.truncate(base_len);
7095            self.spec_ps = p;
7096            self.spec_ress = res;
7097            drafts.clone()
7098        } else if sampling {
7099            let mut p = std::mem::take(&mut self.spec_p);
7100            let mut res = std::mem::take(&mut self.spec_res);
7101            while a < k_spec {
7102                sampler::distribution_into(
7103                    &logits[a * lm_rows..(a + 1) * lm_rows],
7104                    &cfg,
7105                    all_ids,
7106                    &mut self.sampler_scratch,
7107                    self.pool.as_deref(),
7108                    &mut p,
7109                );
7110                match sampler::spec_accept_or_correct(
7111                    &p,
7112                    &self.spec_q[a],
7113                    drafts[a],
7114                    &mut self.rng,
7115                    &mut res,
7116                    self.pool.as_deref(),
7117                ) {
7118                    None => {
7119                        all_ids.push(drafts[a]);
7120                        a += 1;
7121                    }
7122                    Some(c) => {
7123                        forced = Some(c);
7124                        break;
7125                    }
7126                }
7127            }
7128            all_ids.truncate(base_len);
7129            self.spec_p = p;
7130            self.spec_res = res;
7131            // the accepted drafts ARE the verified tokens after inputs 0..a
7132            drafts.clone()
7133        } else if greedy_pen {
7134            // Row i's penalized argmax, penalties over the stream that
7135            // includes the accepted drafts before it — the plain loop's
7136            // exact arithmetic, one pass per row, no working copy.
7137            let mut ids: Vec<u32> = Vec::with_capacity(b);
7138            for i in 0..b {
7139                let t = sampler::argmax_penalized(
7140                    &logits[i * lm_rows..(i + 1) * lm_rows],
7141                    &cfg,
7142                    all_ids,
7143                    &mut self.sampler_scratch,
7144                    self.pool.as_deref(),
7145                );
7146                ids.push(t);
7147                if i < k_spec && t == drafts[i] {
7148                    all_ids.push(t);
7149                } else {
7150                    break;
7151                }
7152            }
7153            all_ids.truncate(base_len);
7154            while a < k_spec && a < ids.len() && ids[a] == drafts[a] {
7155                a += 1;
7156            }
7157            // rows past the first mismatch were never scored; the loop
7158            // top re-samples the last verified row itself.
7159            ids
7160        } else if greedy_dev && dev_ids.len() == b {
7161            let ids = std::mem::take(&mut dev_ids);
7162            while a < k_spec && ids[a] == drafts[a] {
7163                a += 1;
7164            }
7165            ids
7166        } else {
7167            if logits.len() < b * lm_rows {
7168                // the device argmax was asked for and came back short:
7169                // no rows to fall back on — terminal like a failed batch
7170                self.clear_sequence_state();
7171                self.graph_failed
7172                    .store(true, std::sync::atomic::Ordering::Relaxed);
7173                self.cancel
7174                    .store(true, std::sync::atomic::Ordering::Relaxed);
7175                tracing::error!("Metal verify returned neither logits nor argmax ids");
7176                return None;
7177            }
7178            let ids: Vec<u32> = (0..b)
7179                .map(|i| sampler::argmax(&logits[i * lm_rows..(i + 1) * lm_rows]))
7180                .collect();
7181            while a < k_spec && ids[a] == drafts[a] {
7182                a += 1;
7183            }
7184            ids
7185        };
7186        spec_stamp("acc");
7187        if spec_dbg {
7188            eprintln!(
7189                "spec-dbg round: t_next {t_next} drafts {:?} verified {:?} accepted {a}",
7190                drafts, ids
7191            );
7192        }
7193        // CMF_METAL_VERIFY_CHECK=2: the commit oracle — plain-forward the
7194        // a+1 accepted tokens from a snapshot, then diff the replayed GDN
7195        // states and the appended K/V rows against that.
7196        #[cfg(target_os = "macos")]
7197        let commit_ref: Option<(Vec<Vec<f32>>, Vec<(usize, Vec<f32>, Vec<f32>)>)> = if metal_native
7198            && std::env::var("CMF_METAL_VERIFY_CHECK").as_deref() == Ok("2")
7199        {
7200            let snap: Vec<Vec<f32>> = self
7201                .kv_cache
7202                .layers
7203                .iter()
7204                .map(|l| l.linear_state.clone())
7205                .collect();
7206            let attn_lens: Vec<usize> = self.kv_cache.layers.iter().map(|l| l.seq_len).collect();
7207            let toks: Vec<u32> = std::iter::once(t_next)
7208                .chain(drafts.iter().copied())
7209                .collect();
7210            let want_save = self.graph_want_logits;
7211            self.graph_want_logits = false;
7212            for (i, &t) in toks.iter().take(a + 1).enumerate() {
7213                let _ = self.forward_layers(&self.embed_single(t), next_pos + i, None);
7214                let _ = self.graph_logits.take();
7215            }
7216            self.graph_want_logits = want_save;
7217            let plain_states: Vec<Vec<f32>> = self
7218                .kv_cache
7219                .layers
7220                .iter()
7221                .map(|l| l.linear_state.clone())
7222                .collect();
7223            let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7224            let mut rows = Vec::new();
7225            for (li, (l, n0)) in self
7226                .kv_cache
7227                .layers
7228                .iter_mut()
7229                .zip(attn_lens.iter())
7230                .enumerate()
7231            {
7232                let extra = l.seq_len.saturating_sub(*n0);
7233                if extra > 0 {
7234                    let mut kk = Vec::new();
7235                    let mut vv = Vec::new();
7236                    for g in 0..nkv {
7237                        kk.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7238                        vv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7239                    }
7240                    rows.push((li, kk, vv));
7241                    l.truncate_last(extra);
7242                    crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, *n0);
7243                }
7244            }
7245            for (l, st) in self.kv_cache.layers.iter_mut().zip(snap) {
7246                if l.linear_state.len() == st.len() {
7247                    l.linear_state.copy_from_slice(&st);
7248                } else {
7249                    l.linear_state = st;
7250                }
7251            }
7252            Some((plain_states, rows))
7253        } else {
7254            None
7255        };
7256        let warm_off = std::env::var("CMF_SPEC_WARM").is_ok_and(|v| v == "0");
7257        // Metal: the MTP cache cut and the round's warm-up SUBMIT come
7258        // BEFORE the trunk commit, so the warm-up's command buffer is
7259        // queued ahead of the GDN replay (second queue) and its wait
7260        // below no longer sits behind the replay — measured: the warm-up's
7261        // wait grew with the accepted count exactly like the replay does
7262        // (8 ms at a=1, 17 ms at a=3, 25 ms at a=5 for ~2 ms of its own
7263        // work). The replay now overlaps the warm-up's readback, the
7264        // round's return and the next draft chain.
7265        #[cfg(target_os = "macos")]
7266        let mut warm_pending: Option<MetalWarmPending> = None;
7267        #[cfg(target_os = "macos")]
7268        if metal_native {
7269            m.kv.truncate_last(k_spec.saturating_sub(1));
7270            if self.mtp_graph_mode == Some(true) {
7271                // the mirror rows below the cut are the CPU rows: re-point,
7272                // no re-upload
7273                crate::gpu_metal::kv_mirror_set_stored(
7274                    self.mtp_kv_id(),
7275                    Self::MTP_LAYER_BASE,
7276                    m.kv.seq_len,
7277                );
7278                if !warm_off && a > 0 {
7279                    let pairs: Vec<(&[f32], u32)> = (0..a)
7280                        .map(|j| {
7281                            (
7282                                &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size],
7283                                ids[j],
7284                            )
7285                        })
7286                        .collect();
7287                    warm_pending = self.mtp_warm_batch_submit(m, &pairs, next_pos);
7288                }
7289            }
7290            spec_stamp("c.wsub");
7291        }
7292        // a fully-accepted round needs no restore: every input was real.
7293        #[cfg(target_os = "macos")]
7294        if metal_native {
7295            // the Metal verify never wrote its states: the commit replays the
7296            // accepted prefix into the CPU owners and appends the K/V rows
7297            if !self.metal_verify_commit(a) {
7298                self.clear_sequence_state();
7299                self.graph_failed
7300                    .store(true, std::sync::atomic::Ordering::Relaxed);
7301                self.cancel
7302                    .store(true, std::sync::atomic::Ordering::Relaxed);
7303                tracing::error!("Metal verify state/KV handoff failed after admission");
7304                return None;
7305            }
7306            if let Some((plain_states, rows)) = commit_ref {
7307                crate::gpu_metal::queue_fence();
7308                // the commit's replay runs on the second queue: collect it
7309                // before the oracle reads the CPU owners it writes into
7310                let _ = crate::gpu_metal::wait_replay();
7311                let (nkv, hd) = (self.num_kv_heads, self.head_dim);
7312                let mut worst_s = 0f32;
7313                let mut worst_li = 0usize;
7314                for (li, (l, ps)) in self.kv_cache.layers.iter().zip(&plain_states).enumerate() {
7315                    if l.linear_state.len() != ps.len() || ps.is_empty() {
7316                        continue;
7317                    }
7318                    let d = l
7319                        .linear_state
7320                        .iter()
7321                        .zip(ps)
7322                        .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7323                    let n = ps.iter().fold(0f32, |m, y| m.max(y.abs()));
7324                    let rel = d / n.max(1e-6);
7325                    if rel > worst_s {
7326                        worst_s = rel;
7327                        worst_li = li;
7328                    }
7329                }
7330                let mut worst_k = 0f32;
7331                for (li, kk, vv) in &rows {
7332                    let l = &self.kv_cache.layers[*li];
7333                    let n0 = l.seq_len - (kk.len() / (nkv * hd));
7334                    let mut ck = Vec::new();
7335                    let mut cv = Vec::new();
7336                    for g in 0..nkv {
7337                        ck.extend_from_slice(&l.head_keys(g)[n0 * hd..]);
7338                        cv.extend_from_slice(&l.head_values(g)[n0 * hd..]);
7339                    }
7340                    if ck.len() == kk.len() {
7341                        let dk = ck
7342                            .iter()
7343                            .zip(kk)
7344                            .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7345                        let dv = cv
7346                            .iter()
7347                            .zip(vv)
7348                            .fold(0f32, |m, (x, y)| m.max((x - y).abs()));
7349                        worst_k = worst_k.max(dk).max(dv);
7350                    } else {
7351                        eprintln!(
7352                            "commit-check L{li}: kv row count mismatch {} vs {}",
7353                            ck.len(),
7354                            kk.len()
7355                        );
7356                    }
7357                }
7358                eprintln!(
7359                    "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}"
7360                );
7361            }
7362        }
7363        if !metal_native && a + 1 < b {
7364            let expected_gdn_layers = self.graph_gdn_layer_count();
7365            if expected_gdn_layers > 0
7366                && !crate::gpu::gdn_spec_restore(self.graph_kv_id, a, next_pos, expected_gdn_layers)
7367            {
7368                self.clear_sequence_state();
7369                self.graph_failed
7370                    .store(true, std::sync::atomic::Ordering::Relaxed);
7371                self.cancel
7372                    .store(true, std::sync::atomic::Ordering::Relaxed);
7373                tracing::error!("GDN speculative restore failed after verify");
7374                return None;
7375            }
7376        }
7377        if !metal_native && !self.rewind_trunk_graph_mirrors(next_pos + a + 1) {
7378            // The verify graph committed the full batch, but one of its
7379            // persistent Full-attention mirrors could not be re-pointed to
7380            // the accepted prefix.  Treat that as terminal state failure;
7381            // an exact CPU fallback would otherwise consume stale GDN/KV.
7382            self.clear_sequence_state();
7383            self.graph_failed
7384                .store(true, std::sync::atomic::Ordering::Relaxed);
7385            self.cancel
7386                .store(true, std::sync::atomic::Ordering::Relaxed);
7387            tracing::error!("trunk graph KV rewind failed after speculative verify");
7388            return None;
7389        }
7390        *accepted += a;
7391        // MTP cache: keep the first draft row (its inputs were real), drop
7392        // the chain's, then append the verified pairs the round produced.
7393        // Each of those is a whole MTP block on the per-op path and they
7394        // cost 5.8 ms of a 69 ms round at k=3 — a third of what the
7395        // round's own draft costs. PRICED, and they earn it: skipping
7396        // them (`CMF_SPEC_WARM=0`) drops acceptance from 89% to 81% at
7397        // k=3 and 85% to 74% at k=4, and the tok/s goes nowhere at k=3
7398        // (50.3 against 50.5) and backwards at k=4 (48.1 against 50.1).
7399        // The knob stays so the next person can re-price it after the
7400        // warms are batched instead of assuming either way.
7401        if !metal_native {
7402            // (Metal cut its MTP cache before the trunk commit, above)
7403            m.kv.truncate_last(k_spec.saturating_sub(1));
7404        }
7405        spec_stamp("c.trunc");
7406        if !metal_native
7407            && self.mtp_graph_mode == Some(true)
7408            && !self.rewind_mtp_graph_mirror(next_pos)
7409        {
7410            // The graph draft was admitted, so inability to move its cursor
7411            // back to the real anchor is a state failure, not a capability
7412            // refusal.  Do not warm or continue with a stale mirror.
7413            self.clear_sequence_state();
7414            self.graph_failed
7415                .store(true, std::sync::atomic::Ordering::Relaxed);
7416            self.cancel
7417                .store(true, std::sync::atomic::Ordering::Relaxed);
7418            tracing::error!("MTP graph mirror rewind failed after verify commit");
7419            return None;
7420        }
7421        if !warm_off && a > 0 {
7422            // Graph arm: all accepted pairs in ONE batched run over the
7423            // MTP block; the token graph one by one if the batch declines.
7424            let mut warmed = false;
7425            #[cfg(target_os = "macos")]
7426            if metal_native && self.mtp_graph_mode == Some(true) {
7427                // the batched warm-up was submitted before the trunk
7428                // commit: collect it here; one by one on the token graph
7429                // if it declined (or failed)
7430                warmed = match warm_pending.take() {
7431                    Some(p) => self.mtp_warm_batch_finish(m, p),
7432                    None => false,
7433                };
7434                if !warmed {
7435                    warmed = true;
7436                    for j in 0..a {
7437                        let row =
7438                            hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec();
7439                        if self
7440                            .mtp_step_metal(m, &row, ids[j], next_pos + j, false)
7441                            .is_none()
7442                        {
7443                            warmed = false;
7444                            break;
7445                        }
7446                    }
7447                }
7448            }
7449            if !warmed && self.mtp_graph_mode != Some(false) && !metal_native {
7450                let rows: Vec<Vec<f32>> = (0..a)
7451                    .map(|j| hiddens[j * self.hidden_size..(j + 1) * self.hidden_size].to_vec())
7452                    .collect();
7453                let pairs: Vec<(&[f32], u32)> = rows
7454                    .iter()
7455                    .zip(ids.iter())
7456                    .map(|(r, &t)| (r.as_slice(), t))
7457                    .collect();
7458                match self.mtp_warm_prefill_pairs(m, &pairs, next_pos) {
7459                    Ok(()) => warmed = true,
7460                    Err(err) => {
7461                        // A warm-up failure after graph admission cannot
7462                        // fall back to `mtp_warm`: the detached CPU cache is
7463                        // not authoritative for the device mirror.  Mark it
7464                        // terminal so the generation caller clears state and
7465                        // returns instead of drafting from stale attention.
7466                        tracing::error!("{err}");
7467                        self.clear_sequence_state();
7468                        self.graph_failed
7469                            .store(true, std::sync::atomic::Ordering::Relaxed);
7470                        self.cancel
7471                            .store(true, std::sync::atomic::Ordering::Relaxed);
7472                        return None;
7473                    }
7474                }
7475            }
7476            if !warmed {
7477                for j in 0..a {
7478                    let row = &hiddens[j * self.hidden_size..(j + 1) * self.hidden_size];
7479                    let row = row.to_vec();
7480                    self.mtp_warm(m, &row, ids[j], next_pos + j);
7481                }
7482            }
7483        }
7484        // The sampler's contract: logits of the LAST verified position —
7485        // unless a rejected draft already drew the correction, in which
7486        // case the loop top commits that token and samples nothing.
7487        spec_stamp("c.warm");
7488        if let Some(c) = forced {
7489            self.spec_forced = Some(c);
7490            self.graph_logits = None;
7491        } else if greedy_dev && logits.is_empty() {
7492            // the row's argmax IS the token the loop top would pick from
7493            // it (plain greedy, no penalties): commit it as forced
7494            self.spec_forced = Some(ids[a]);
7495            self.graph_logits = None;
7496        } else {
7497            let mut row = logits[a * lm_rows..(a + 1) * lm_rows].to_vec();
7498            row.resize(self.vocab_size, 0.0);
7499            if let Some(c) = self.final_softcap {
7500                for l in row.iter_mut() {
7501                    *l = c * (*l / c).tanh();
7502                }
7503            }
7504            self.graph_logits = Some(row);
7505        }
7506        let new_hidden = hiddens[a * self.hidden_size..(a + 1) * self.hidden_size].to_vec();
7507        spec_stamp("c.row");
7508        // Three phases, not two. The round's wall clock was 4 ms longer
7509        // than draft+verify and the difference had nowhere to be seen:
7510        // the accepted prefix re-runs the MTP block once per token to
7511        // keep the draft head's attention cache warm, and the GDN state
7512        // rolls back on any rejection. Both live here, after the verify.
7513        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7514            let end = subs();
7515            eprintln!(
7516                "spec-round: draft {:.1} ms/{} sub | verify {:.1} ms/{} sub | \
7517                 commit {:.1} ms/{} sub (accepted {a} of {k_spec}, full-head streak {})",
7518                t_draft.as_secs_f64() * 1e3,
7519                sub_draft - sub0,
7520                (t_verify - t_draft).as_secs_f64() * 1e3,
7521                sub_verify - sub_draft,
7522                (t_round.elapsed() - t_verify).as_secs_f64() * 1e3,
7523                end - sub_verify,
7524                self.draft_full_streak,
7525            );
7526        }
7527        // Native Metal's verify tile is flat in b (eight rows for the price
7528        // of one), so a shorter round only forfeits tokens — measured on
7529        // the M4: an essay round at k=2 still verified in 260 ms. The
7530        // adaptation is for cards whose verify grows with the rows.
7531        if k_env.is_none() && !metal_native && !k_capped {
7532            // Slow average and a wide band: a fast one oscillated 2↔3 on
7533            // an essay every other round (measured), which forfeits the
7534            // draft it just paid for.
7535            let f = a as f32 / k_spec.max(1) as f32;
7536            self.spec_acc_ewma += 0.2 * (f - self.spec_acc_ewma);
7537            let mut k_next = k_spec;
7538            if self.spec_acc_ewma >= 0.75 && k_spec < k_max {
7539                k_next = k_spec + 1;
7540            } else if self.spec_acc_ewma < 0.4 && k_spec > 2 {
7541                k_next = k_spec - 1;
7542            }
7543            if k_next != k_spec {
7544                self.spec_acc_ewma = 0.6;
7545                if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
7546                    eprintln!("spec-k: {k_spec} → {k_next}");
7547                }
7548            }
7549            self.spec_k_adapt = Some(k_next);
7550        }
7551        spec_stamp("end");
7552        Some((drafts[..a].to_vec(), next_pos + a + 1, new_hidden))
7553    }
7554
7555    /// Micro-benchmark: two single-position forwards vs one fused pair
7556    /// from the current cache state (KV rewound after each probe).
7557    /// Returns (two_singles_ms, fused_pair_ms) per probe, or the (0, 0)
7558    /// sentinel when this model has no pair path to measure — the same
7559    /// answer the o1 arm gives, and the bench prints it the same way.
7560    /// (An architecture that loads its own layers leaves `weights.layers`
7561    /// empty; walking it here was an index panic, found by `bench` on
7562    /// deepseek_v4.)
7563    pub fn measure_pair_fusion(&mut self, iters: usize) -> (f64, f64) {
7564        if !self.pair_supported() {
7565            return (0.0, 0.0);
7566        }
7567        // This is a host-side pair micro-benchmark. It truncates the host KV
7568        // after every probe, so letting the whole-token graph participate
7569        // would leave its device GDN/KV mirror ahead of the next probe and
7570        // poison the process-wide graph verdict before the real generation
7571        // benchmark starts. Keep the existing per-op/GPU arithmetic while
7572        // suppressing only the stateful token graph for this measurement.
7573        let graph_env = std::env::var_os("CMF_GPU_WGPU_GRAPH");
7574        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
7575        let emb1 = self.embed_single(1);
7576        let emb2 = self.embed_single(2);
7577        let pos = self.kv_cache.seq_len();
7578
7579        let t0 = std::time::Instant::now();
7580        for _ in 0..iters {
7581            let _ = self.forward_layers(&emb1, pos, None);
7582            let _ = self.forward_layers(&emb2, pos + 1, None);
7583            for l in &mut self.kv_cache.layers {
7584                l.truncate_last(2);
7585            }
7586        }
7587        let singles_ms = t0.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7588
7589        let t1 = std::time::Instant::now();
7590        for _ in 0..iters {
7591            let _ = self.forward_pair(&emb1, &emb2, pos);
7592            for l in &mut self.kv_cache.layers {
7593                l.truncate_last(2);
7594            }
7595        }
7596        let pair_ms = t1.elapsed().as_secs_f64() * 1000.0 / iters as f64;
7597        match graph_env {
7598            Some(value) => unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", value) },
7599            None => unsafe { std::env::remove_var("CMF_GPU_WGPU_GRAPH") },
7600        }
7601        (singles_ms, pair_ms)
7602    }
7603
7604    /// Fused two-position forward: weight rows are streamed from memory
7605    /// once per layer for both positions. Full layers → fused GQA pair;
7606    /// linear layers → vmf_phase pair (lane 2 state is tentative in the
7607    /// per-layer scratch until the draft is accepted).
7608    /// Whether the fused two-position path covers every layer kind in
7609    /// this model. MLA and KDA run per position (their pair arms are
7610    /// unreachable); the seq prefill falls back to singles for them.
7611    fn pair_supported(&self) -> bool {
7612        // An EMPTY layer stack means the architecture loaded its own and
7613        // this path has nothing to walk. Checking that directly, rather
7614        // than naming each such architecture, is what makes the guard hold
7615        // for the next one: `any()` over no layers is false, so a
7616        // feature-by-feature test says "supported" for a model that has no
7617        // layers here at all.
7618        !self.weights.layers.is_empty()
7619            && self.g3n.is_none()
7620            && !self
7621                .weights
7622                .layers
7623                .iter()
7624                .any(|lw| matches!(&lw.attn, AttnKind::Mla(_) | AttnKind::Kda(_)))
7625    }
7626
7627    fn forward_pair(
7628        &mut self,
7629        emb1: &[f32],
7630        emb2: &[f32],
7631        position: usize,
7632    ) -> (Vec<f32>, Vec<f32>) {
7633        // A two-token prompt starts here, not in the layer walk: decide the
7634        // MiMo placement before the pair's per-op MoE uploads any expert.
7635        self.mimo_moe_prepare();
7636        let mut h1 = emb1.to_vec();
7637        let mut h2 = emb2.to_vec();
7638        let (_nkv, _hd, hs, _rd, eps) = (
7639            self.num_kv_heads,
7640            self.head_dim,
7641            self.hidden_size,
7642            self.rotary_dim,
7643            self.rms_eps,
7644        );
7645        let pool = self.pool.clone();
7646
7647        for li in 0..self.num_layers {
7648            let lw = &self.weights.layers[self.phys_layer(li)];
7649            // Norms into pipeline scratch (4 allocs/layer on the MTP
7650            // decode hot path before this).
7651            inference::rms_norm_into(
7652                &h1,
7653                &lw.input_norm,
7654                self.rms_eps,
7655                self.norm_style,
7656                &mut self.ws.n1,
7657            );
7658            inference::rms_norm_into(
7659                &h2,
7660                &lw.input_norm,
7661                self.rms_eps,
7662                self.norm_style,
7663                &mut self.ws.n2,
7664            );
7665
7666            let (a1, a2) = match &lw.attn {
7667                AttnKind::Mla(_) => unreachable!("MLA has no MTP/pair path"),
7668                AttnKind::Kda(_) => unreachable!("KDA has no MTP/pair path"),
7669                AttnKind::Bounded(w) => {
7670                    // Two sequential positions of the bounded operator
7671                    // (the ring's causal order is the pair's order).
7672                    let rope = self
7673                        .bounded_rope
7674                        .clone()
7675                        .expect("bounded layer without an installed rotation table");
7676                    let cfg = crate::bounded::BoundedAttnCfg {
7677                        num_heads: self.num_heads,
7678                        num_kv_heads: self.num_kv_heads,
7679                        head_dim: self.head_dim,
7680                        hidden_size: hs,
7681                        scale: self.attn_scale,
7682                        rope: &rope,
7683                        pool: pool.as_deref(),
7684                    };
7685                    let a1 = crate::bounded::bounded_attention(
7686                        &self.ws.n1,
7687                        w,
7688                        &mut self.kv_cache.layers[li],
7689                        &cfg,
7690                    );
7691                    let a2 = crate::bounded::bounded_attention(
7692                        &self.ws.n2,
7693                        w,
7694                        &mut self.kv_cache.layers[li],
7695                        &cfg,
7696                    );
7697                    (a1, a2)
7698                }
7699                AttnKind::Linear(w) => {
7700                    let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
7701                    let layer = &mut self.kv_cache.layers[li];
7702                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7703                    vmf_phase_pair(
7704                        &self.ws.n1,
7705                        &self.ws.n2,
7706                        w,
7707                        &cfg,
7708                        state,
7709                        scratch,
7710                        self.pool.as_deref(),
7711                    )
7712                }
7713                AttnKind::LinearGdn(w) => {
7714                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
7715                    let layer = &mut self.kv_cache.layers[li];
7716                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7717                    gdn_pair(
7718                        &self.ws.n1,
7719                        &self.ws.n2,
7720                        w,
7721                        &cfg,
7722                        state,
7723                        scratch,
7724                        self.pool.as_deref(),
7725                    )
7726                }
7727                AttnKind::ShortConv(w) => {
7728                    let cfg = self
7729                        .short_conv_cfg
7730                        .expect("short-conv layer without short_conv_cfg");
7731                    let layer = &mut self.kv_cache.layers[li];
7732                    let (state, scratch) = (&mut layer.linear_state, &mut layer.linear_scratch);
7733                    short_conv_pair(
7734                        &self.ws.n1,
7735                        &self.ws.n2,
7736                        w,
7737                        &cfg,
7738                        state,
7739                        scratch,
7740                        self.pool.as_deref(),
7741                    )
7742                }
7743                AttnKind::Full {
7744                    wq,
7745                    wk,
7746                    wv,
7747                    wo,
7748                    q_norm,
7749                    k_norm,
7750                    output_gate,
7751                    softplus_gate,
7752                    bias,
7753                } => {
7754                    let inv_freq_l = self.layer_inv_freq(li);
7755                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
7756                    let cfg = QwenAttnCfg {
7757                        num_heads: self.layer_num_heads(li),
7758                        num_kv_heads: nkv_l,
7759                        head_dim: hd_l,
7760                        hidden_size: hs,
7761                        position,
7762                        inv_freq: &inv_freq_l,
7763                        rotary_dim: rd_l,
7764                        scale: self.attn_scale,
7765                        softcap: self.attn_softcap,
7766                        window: self.layer_window(li),
7767                        v_norm: self.attn_v_norm,
7768                        qk_norm_after_rope: self.qk_norm_after_rope,
7769                        gate_sigmoid: self.proj_gate_sigmoid,
7770                        q_norm: q_norm.as_deref(),
7771                        k_norm: k_norm.as_deref(),
7772                        output_gate: *output_gate,
7773                        softplus_gate: softplus_gate
7774                            .as_ref()
7775                            .map(|(gate, per_head)| (gate, *per_head)),
7776                        rope_scale: self.layer_rope_scale(li),
7777                        bias: bias
7778                            .as_ref()
7779                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
7780                        rms_eps: eps,
7781                        norm_style: self.norm_style,
7782                        pool: pool.as_deref(),
7783                        v_head_dim: self.layer_v_dim(li),
7784                    };
7785                    attention::qwen_attention_pair(
7786                        &self.ws.n1,
7787                        &self.ws.n2,
7788                        wq,
7789                        wk,
7790                        wv,
7791                        wo,
7792                        &mut self.kv_cache.layers[li],
7793                        &cfg,
7794                    )
7795                }
7796            };
7797            let (a1, a2) = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
7798                Some(w) => (
7799                    inference::rms_norm(&a1, w, self.rms_eps, self.norm_style),
7800                    inference::rms_norm(&a2, w, self.rms_eps, self.norm_style),
7801                ),
7802                None => (a1, a2),
7803            };
7804            for i in 0..self.hidden_size {
7805                h1[i] += a1[i];
7806                h2[i] += a2[i];
7807            }
7808            let (mut a1, mut a2) = (a1, a2);
7809            attention::recycle_buf(&mut a1);
7810            attention::recycle_buf(&mut a2);
7811
7812            let lw = &self.weights.layers[self.phys_layer(li)];
7813            inference::rms_norm_into(
7814                &h1,
7815                &lw.post_norm,
7816                self.rms_eps,
7817                self.norm_style,
7818                &mut self.ws.p1,
7819            );
7820            inference::rms_norm_into(
7821                &h2,
7822                &lw.post_norm,
7823                self.rms_eps,
7824                self.norm_style,
7825                &mut self.ws.p2,
7826            );
7827            let (f1, f2) = match &lw.ffn {
7828                // Dual-branch layers need the raw residuals — run the
7829                // two positions through the same fn decode uses.
7830                FfnKind::DenseMoe(dm) => (
7831                    dense_moe_ffn(
7832                        dm,
7833                        &self.ws.p1,
7834                        &h1,
7835                        self.rms_eps,
7836                        self.norm_style,
7837                        self.pool.as_deref(),
7838                    ),
7839                    dense_moe_ffn(
7840                        dm,
7841                        &self.ws.p2,
7842                        &h2,
7843                        self.rms_eps,
7844                        self.norm_style,
7845                        self.pool.as_deref(),
7846                    ),
7847                ),
7848                FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => (
7849                    moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p1, self.pool.as_deref()),
7850                    moe_ffn_banked(&mut self.mimo_moe, li, m, &self.ws.p2, self.pool.as_deref()),
7851                ),
7852                _ => ffn_forward_pair(
7853                    &lw.ffn,
7854                    &self.ws.p1,
7855                    &self.ws.p2,
7856                    self.pool.as_deref(),
7857                    None,
7858                ),
7859            };
7860            let (f1, f2) = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
7861                Some(w) => (
7862                    inference::rms_norm(&f1, w, self.rms_eps, self.norm_style),
7863                    inference::rms_norm(&f2, w, self.rms_eps, self.norm_style),
7864                ),
7865                None => (f1, f2),
7866            };
7867            for i in 0..self.hidden_size {
7868                h1[i] += f1[i];
7869                h2[i] += f2[i];
7870            }
7871            let (mut f1, mut f2) = (f1, f2);
7872            attention::recycle_buf(&mut f1);
7873            attention::recycle_buf(&mut f2);
7874            if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
7875                for i in 0..self.hidden_size {
7876                    h1[i] *= sc;
7877                    h2[i] *= sc;
7878                }
7879            }
7880            // Looped Transformer: apply final norm at the end of each loop iteration.
7881            if self.is_loop_end(li) && li + 1 < self.num_layers {
7882                h1 = inference::rms_norm(
7883                    &h1,
7884                    &self.weights.final_norm,
7885                    self.rms_eps,
7886                    self.norm_style,
7887                );
7888                h2 = inference::rms_norm(
7889                    &h2,
7890                    &self.weights.final_norm,
7891                    self.rms_eps,
7892                    self.norm_style,
7893                );
7894            }
7895        }
7896        // Real O(1) prefill pairs may also carry tentative lane-2 recurrent
7897        // state. Commit it before publishing the transition epoch so the
7898        // next serial/device row cannot observe a new attention epoch with an
7899        // old GDN state. Speculative pairs run only when O(1) is inactive and
7900        // retain their existing caller-controlled commit/rollback semantics.
7901        if self.o1_active() {
7902            self.commit_linear_scratch();
7903        }
7904        self.o1_progress();
7905        self.swa_trim_tails();
7906        (h1, h2)
7907    }
7908
7909    /// Commit lane-2 linear states after an accepted draft.
7910    fn commit_linear_scratch(&mut self) {
7911        for layer in &mut self.kv_cache.layers {
7912            if !layer.linear_scratch.is_empty() {
7913                std::mem::swap(&mut layer.linear_state, &mut layer.linear_scratch);
7914                layer.linear_scratch.clear();
7915            }
7916        }
7917    }
7918
7919    /// Forward a full id sequence from a fresh cache and return the
7920    /// logits after the last position (golden-parity harness, bench).
7921    pub fn forward_ids(
7922        &mut self,
7923        ids: &[u32],
7924        task_mask: Option<&TaskMask>,
7925    ) -> Result<Vec<f32>, String> {
7926        #[cfg(target_os = "macos")]
7927        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
7928        if ids.is_empty() {
7929            return Err("empty id sequence".to_string());
7930        }
7931        self.clear_sequence_state();
7932        self.check_forward_graph("forward_ids setup", 0)?;
7933        if task_mask.is_none() {
7934            self.o1_begin();
7935        }
7936        let mut hidden = vec![0.0f32; self.hidden_size];
7937        let mut pos = 0usize;
7938        if let Some(b) = &mut self.dsv41 {
7939            let pool = self.pool.clone();
7940            let mut logits = Vec::new();
7941            crate::dsv41::forward_chunk(
7942                &b.0,
7943                &b.1,
7944                &b.2,
7945                &mut b.3,
7946                ids,
7947                0,
7948                pool.as_deref(),
7949                &mut logits,
7950            );
7951            if let Err(err) = self.o1_seal_checked() {
7952                self.clear_sequence_state();
7953                return Err(err);
7954            }
7955            return Ok(logits);
7956        }
7957        // Same routing predicate generation uses. Two reasons it must be
7958        // the same one: (1) a GDN hybrid's recurrent state is GPU-
7959        // resident, and a batched CPU prefill would build it on the host
7960        // only — decode then reads buffers the prefill never wrote;
7961        // (2) bench times THIS function and calls the result "prefill",
7962        // so a different path here reports a number production never
7963        // sees (W2 on 2×5090: 8.7 tok/s reported against 125 real).
7964        if self.can_prefill_batched() && !self.graph_prefill_preferred() && ids.len() > 2 {
7965            // prefill-GEMM in chunks; only the last position's hidden is
7966            // needed. (o1-compatible: the batch path attends per position
7967            // through qwen_attention, which carries the collection hook.)
7968            let chunk = self.prefill_chunk();
7969            let hs = self.hidden_size;
7970            while pos < ids.len() {
7971                let end = (pos + chunk).min(ids.len());
7972                let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
7973                self.check_forward_graph("forward_ids batched prefill", end - 1)?;
7974                hidden.copy_from_slice(&hb[(end - pos - 1) * hs..]);
7975                pos = end;
7976            }
7977        }
7978        // Same guards as generation's prefill — INCLUDING the graph one.
7979        // The CPU pair walk was intercepting positions that the resident
7980        // token graph would have run itself: on a GDN hybrid over wgpu
7981        // that is 89 ms of host forward against 7 ms of device submit,
7982        // and it made prefill look 12× slower than it is (W2 on an RTX
7983        // 5090, ctx 512: 11.2 tok/s with the walk, 136.6 without).
7984        // CMF_PAIR=0 opts out; a model whose layers live outside
7985        // `weights.layers` has no pair walk to take.
7986        if task_mask.is_none()
7987            && !self.graph_prefill_preferred()
7988            && !std::env::var("CMF_PAIR").is_ok_and(|v| v == "0")
7989            && self.pair_supported()
7990        {
7991            while pos + 1 < ids.len() {
7992                let e1 = self.embed_single(ids[pos]);
7993                let e2 = self.embed_single(ids[pos + 1]);
7994                let (_, h2) = self.forward_pair(&e1, &e2, pos);
7995                self.check_forward_graph("forward_ids pair", pos + 1)?;
7996                self.commit_linear_scratch();
7997                hidden = h2;
7998                pos += 2;
7999            }
8000        }
8001        // Resident Embryo graph: the prompt in chunks of one submit each
8002        // (the same device state and logits as the per-position walk).
8003        if task_mask.is_none() && pos == 0 && ids.len() > 1 {
8004            if let Some(lg) = self.embryo_prefill_chunked(ids, 0) {
8005                self.graph_logits = Some(lg);
8006                hidden = vec![0.0; self.hidden_size];
8007                pos = ids.len();
8008            }
8009        }
8010        while pos < ids.len() {
8011            hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
8012            self.check_forward_graph("forward_ids", pos)?;
8013            pos += 1;
8014        }
8015        if let Some(logits) = self.graph_logits.take() {
8016            // Resident stacks already applied final norm and their head in
8017            // the same submit; do not run a second norm/head over the zero
8018            // hidden sentinel returned by forward_layers_span.
8019            if let Err(err) = self.o1_seal_checked() {
8020                self.clear_sequence_state();
8021                return Err(err);
8022            }
8023            return Ok(logits);
8024        }
8025        // Harness contract: after forward_ids the cache is decode-ready —
8026        // under o1 that means sealed (bench measures the seal as part of
8027        // prefill, honestly).
8028        if let Err(err) = self.o1_seal_checked() {
8029            self.clear_sequence_state();
8030            return Err(err);
8031        }
8032        let normed = inference::rms_norm(
8033            &hidden,
8034            &self.weights.final_norm,
8035            self.rms_eps,
8036            self.norm_style,
8037        );
8038        Ok(self.lm_head_forward(&normed))
8039    }
8040
8041    /// Run the V4.1 stack one token at a time and retain logits for every
8042    /// position. This is a diagnostic surface for comparing a converted
8043    /// checkpoint with a tokenwise reference implementation.
8044    #[doc(hidden)]
8045    pub fn dsv41_serial_logits(&mut self, ids: &[u32]) -> Result<Vec<Vec<f32>>, String> {
8046        #[cfg(target_os = "macos")]
8047        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8048        if ids.is_empty() {
8049            return Err("empty id sequence".to_string());
8050        }
8051        self.clear_sequence_state();
8052        self.dsv41
8053            .as_ref()
8054            .ok_or_else(|| "dsv41 serial logits require a DeepSeek-V4.1 model".to_string())?;
8055        self.o1_begin();
8056        let rows = {
8057            let pool = self.pool.clone();
8058            let b = self
8059                .dsv41
8060                .as_mut()
8061                .expect("dsv41 checked above; state cannot change during forward");
8062            let mut rows = Vec::with_capacity(ids.len());
8063            for (position, &id) in ids.iter().enumerate() {
8064                let mut logits = Vec::new();
8065                crate::dsv41::forward_token(
8066                    &b.0,
8067                    &b.1,
8068                    &b.2,
8069                    &mut b.3,
8070                    id,
8071                    position,
8072                    pool.as_deref(),
8073                    &mut logits,
8074                );
8075                rows.push(logits);
8076            }
8077            rows
8078        };
8079        self.o1_seal();
8080        Ok(rows)
8081    }
8082
8083    /// Teacher-forced perplexity over a token sequence (phase-C gate:
8084    /// honest quant comparisons instead of prompt vibes).
8085    ///
8086    /// Attention is EXACT even on a model whose layers are flagged for
8087    /// the O(1) kernel — scoring the backbone is the default on purpose
8088    /// (it is the yardstick). `nll_ids_o1` scores the CONVERTED model.
8089    pub fn ppl_ids(&mut self, ids: &[u32]) -> Result<f64, String> {
8090        let (nll, cnt) = self.nll_ids_from(ids, 0)?;
8091        Ok((nll / cnt.max(1) as f64).exp())
8092    }
8093
8094    /// DTG-MA calibration pass (Patent 2): run `ids` through the model
8095    /// (CPU path, per position) and return each layer's per-neuron
8096    /// activation mass Σ|silu(gate)·up| — the statistic the task-guided
8097    /// FFN mask is derived from.
8098    pub fn probe_ffn_mass(&mut self, ids: &[u32]) -> Vec<Vec<f64>> {
8099        self.clear_sequence_state();
8100        FFN_PROBE.with(|p| {
8101            *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
8102        });
8103        crate::gpu::cpu_scope(|| {
8104            for (pos, &id) in ids.iter().enumerate() {
8105                let emb = self.embed_single(id);
8106                let _ = self.forward_layers(&emb, pos, None);
8107            }
8108        });
8109        self.clear_sequence_state();
8110        FFN_PROBE
8111            .with(|p| p.borrow_mut().take())
8112            .unwrap_or_default()
8113    }
8114
8115    /// `probe_ffn_mass` over the BATCHED prefill: same accumulator, one
8116    /// sweep instead of one forward per token. What makes the statistic
8117    /// affordable on a 27B.
8118    pub fn probe_ffn_mass_batch(&mut self, ids: &[u32]) -> Result<Vec<Vec<f64>>, String> {
8119        if let Err(err) = self.nll_begin() {
8120            // A recorder can be left by a caller that was interrupted before
8121            // this request entered its scoring block.  Consume it even when
8122            // the preflight failure prevents initialization of a new one.
8123            let _ = FFN_PROBE.with(|p| p.borrow_mut().take());
8124            self.nll_end();
8125            return Err(err);
8126        }
8127        FFN_PROBE.with(|p| {
8128            *p.borrow_mut() = Some(vec![vec![0f64; self.intermediate_size]; self.num_layers]);
8129        });
8130        let result: Result<(), String> = (|| {
8131            for chunk in ids.chunks(256) {
8132                if chunk.len() < 2 {
8133                    continue;
8134                }
8135                self.nll_ids_masked(chunk, 0, None)?;
8136            }
8137            Ok(())
8138        })();
8139        self.nll_end();
8140        let probe = FFN_PROBE
8141            .with(|p| p.borrow_mut().take())
8142            .unwrap_or_default();
8143        match result {
8144            Ok(()) => Ok(probe),
8145            Err(err) => {
8146                drop(probe);
8147                Err(err)
8148            }
8149        }
8150    }
8151
8152    /// Teacher-forced PPL with a task mask active (sparse execution) —
8153    /// the quality gate for a DTG-MA-masked skill. Sequential per
8154    /// position: the batched prefill path is dense-only.
8155    pub fn ppl_ids_masked(&mut self, ids: &[u32], mask: &TaskMask) -> Result<f64, String> {
8156        self.nll_begin()?;
8157        let result: Result<f64, String> = (|| {
8158            let mut nll = 0f64;
8159            let mut cnt = 0usize;
8160            let mut hidden = vec![0f32; self.hidden_size];
8161            for (pos, &id) in ids.iter().enumerate() {
8162                if pos > 0 {
8163                    inference::rms_norm_into(
8164                        &hidden,
8165                        &self.weights.final_norm,
8166                        self.rms_eps,
8167                        self.norm_style,
8168                        &mut self.ws.n1,
8169                    );
8170                    let mut logits = self.lm_head_forward(&self.ws.n1);
8171                    let max = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
8172                    let sum: f64 = logits.iter().map(|&v| ((v - max) as f64).exp()).sum();
8173                    let p = ((logits[id as usize] - max) as f64).exp() / sum.max(1e-300);
8174                    nll -= p.max(1e-300).ln();
8175                    cnt += 1;
8176                    attention::recycle_buf(&mut logits);
8177                }
8178                let emb = self.embed_single(id);
8179                hidden = self.forward_layers(&emb, pos, Some(mask));
8180                self.nll_check_graph("masked serial forward", pos)?;
8181                // Consume a possible graph logits side channel before the
8182                // next row.  Masked scoring normally disables that route,
8183                // but stale channel state must never survive a request.
8184                let _ = self.graph_logits.take();
8185            }
8186            Ok((nll / cnt.max(1) as f64).exp())
8187        })();
8188        self.nll_end();
8189        result
8190    }
8191
8192    /// Teacher-forced NLL sum + scored-token count over positions
8193    /// `start..len-1`, attention EXACT. Positions below `start` still
8194    /// run — they are the context — they are just not scored, so this
8195    /// pairs with `nll_ids_o1(ids, start)` over the very same tokens.
8196    ///
8197    /// Returning (nll, cnt) rather than a ppl is what lets a windowed
8198    /// caller combine windows before the exp, so every scored token
8199    /// weighs the same regardless of how the windows are cut.
8200    /// `nll_ids_from` with a task mask held active at every position.
8201    ///
8202    /// The batched prefill path does not thread masks, so this walks the
8203    /// per-position forward — slower, but it scores the file exactly the
8204    /// way `run --task` will serve it, which is the point of the gate
8205    /// that calls it. With `None` it defers to the fast path.
8206    /// Masked scoring rides the SAME batched sweep as unmasked scoring —
8207    /// the masked-inference fast path: `prefill_batch_masked` lands the
8208    /// per-visit FFN rows on the activations inside the fused arms. The
8209    /// per-position loop below remains only as the no-batch fallback.
8210    pub fn nll_ids_masked(
8211        &mut self,
8212        ids: &[u32],
8213        start: usize,
8214        task_mask: Option<&TaskMask>,
8215    ) -> Result<(f64, usize), String> {
8216        let task_mask = self.drop_open_mask(task_mask);
8217        self.nll_ids_inner(ids, start, task_mask)
8218    }
8219
8220    pub fn nll_ids_from(&mut self, ids: &[u32], start: usize) -> Result<(f64, usize), String> {
8221        self.nll_ids_inner(ids, start, None)
8222    }
8223
8224    fn nll_ids_inner(
8225        &mut self,
8226        ids: &[u32],
8227        start: usize,
8228        task_mask: Option<&TaskMask>,
8229    ) -> Result<(f64, usize), String> {
8230        self.nll_begin()?;
8231        let result: Result<(f64, usize), String> = (|| {
8232            let mut nll = 0f64;
8233            let mut cnt = 0usize;
8234            // An unmasked quality run with the resident wgpu graph must score
8235            // the same stateful path used by generation.  The layer-major
8236            // GEMM prefill below is a valid CPU/GEMM oracle, but it seeds
8237            // neither the graph's device GDN state nor its device KV mirrors;
8238            // using it here would silently score a different execution.  Keep
8239            // masked scoring on the exact per-position path as before, and
8240            // let the serial arm below drive the graph-aware scorer.
8241            // Only native Metal has a fused graph lm_head contract.  Vulkan
8242            // and other graph backends may expose hidden state without the
8243            // optional logits side channel; preserve their established CPU
8244            // norm/head fallback instead of turning that valid route into a
8245            // hard missing-logits error.
8246            let (graph_quality, fused_head_quality) = nll_graph_policy(
8247                task_mask.is_none(),
8248                self.graph_prefill_preferred(),
8249                crate::gpu::q1_force(),
8250            );
8251            self.graph_head_required = fused_head_quality;
8252            self.graph_want_logits = fused_head_quality;
8253            #[cfg(target_os = "macos")]
8254            if graph_quality && std::env::var("CMF_METAL_BATCH_NLL").as_deref() != Ok("0") {
8255                match self.nll_batch_metal(ids, start) {
8256                    MetalBatchNllOutcome::Completed(nll, count) => {
8257                        return Ok((nll, count));
8258                    }
8259                    MetalBatchNllOutcome::Declined => {}
8260                    MetalBatchNllOutcome::Failed(err) => return Err(err),
8261                }
8262            }
8263            // CMF_NLL_SERIAL=1 scores position by position through the
8264            // decode path (the wgpu token graph where it admits the model)
8265            // — the perplexity gate of a decode kernel, which the batched
8266            // prefill arm below never runs.
8267            let force_serial = std::env::var("CMF_NLL_SERIAL").as_deref() == Ok("1");
8268            if self.can_prefill_batched() && !graph_quality && !force_serial {
8269                // prefill-GEMM: layer-major position chunks, lm_head batched
8270                // (254MB lm_head read once per chunk, not per position).
8271                // The layer chunk is large (grouping positions by MoE experts
8272                // wins with size), lm_head in sub-blocks (logit buffer
8273                // 32×vocab ≈ 32MB instead of 128×).
8274                const CHUNK: usize = 128;
8275                const LM_SUB: usize = 32;
8276                let n = ids.len().saturating_sub(1);
8277                let hs = self.hidden_size;
8278                let rows = self.weights.lm_head.rows();
8279                let mut pos = 0usize;
8280                let state_trace = std::env::var("CMF_STATE_TRACE").is_ok();
8281                while pos < n {
8282                    let end = (pos + CHUNK).min(n);
8283                    let bsz = end - pos;
8284                    let hb = self.prefill_rows(&ids[pos..end], pos, task_mask)?;
8285                    self.nll_check_graph("batched prefill", pos)?;
8286                    if state_trace && end % 256 == 0 {
8287                        self.trace_recurrent_state(end);
8288                    }
8289                    let mut k0 = 0usize;
8290                    while k0 < bsz {
8291                        let k1 = (k0 + LM_SUB).min(bsz);
8292                        let sb = k1 - k0;
8293                        // Sub-block entirely below the scored range: the KV
8294                        // it just built is all this pass needed from it.
8295                        if pos + k1 <= start {
8296                            k0 = k1;
8297                            continue;
8298                        }
8299                        let mut normed = vec![0.0f32; sb * hs];
8300                        for k in 0..sb {
8301                            let r = inference::rms_norm(
8302                                &hb[(k0 + k) * hs..(k0 + k + 1) * hs],
8303                                &self.weights.final_norm,
8304                                self.rms_eps,
8305                                self.norm_style,
8306                            );
8307                            normed[k * hs..(k + 1) * hs].copy_from_slice(&r);
8308                        }
8309                        let mut logits = vec![0.0f32; sb * rows];
8310                        self.weights
8311                            .lm_head
8312                            .matmat(&normed, sb, &mut logits, self.pool.as_deref());
8313                        for k in 0..sb {
8314                            if pos + k0 + k < start {
8315                                continue;
8316                            }
8317                            self.nll_check_graph("batched score row", pos + k0 + k)?;
8318                            let lg = &mut logits[k * rows..k * rows + self.vocab_size.min(rows)];
8319                            if let Some(mu) = self.logit_multiplier {
8320                                for v in lg.iter_mut() {
8321                                    *v *= mu;
8322                                }
8323                            }
8324                            // Gemma-class final-logit soft-capping: the
8325                            // decode paths apply it; scoring must too, or
8326                            // the uncapped softmax misprices every token.
8327                            if let Some(c) = self.final_softcap {
8328                                for v in lg.iter_mut() {
8329                                    *v = c * (*v / c).tanh();
8330                                }
8331                            }
8332                            // Cortiq Embryo hierarchical head: same correction
8333                            // the decode path applies (lm_head_forward).
8334                            if let Some(cm) = self.head_clusters.clone() {
8335                                self.hierarchical_head_logprobs(
8336                                    &normed[k * hs..(k + 1) * hs],
8337                                    &cm,
8338                                    lg,
8339                                );
8340                            }
8341                            let lg = &logits[k * rows..k * rows + self.vocab_size.min(rows)];
8342                            let target = ids[pos + k0 + k + 1] as usize;
8343                            let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8344                            let lse: f64 = lg
8345                                .iter()
8346                                .map(|&v| ((v - max) as f64).exp())
8347                                .sum::<f64>()
8348                                .ln()
8349                                + max as f64;
8350                            nll += lse - lg[target] as f64;
8351                            cnt += 1;
8352                            if std::env::var("CMF_PPL_TRACE").is_ok() {
8353                                let top = lg
8354                                    .iter()
8355                                    .enumerate()
8356                                    .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8357                                    .map(|(i, _)| i)
8358                                    .unwrap_or(0);
8359                                eprintln!(
8360                                    "BTRACE pos {} target {} nll {:.4} top {} lg_t {:.3} lg_top {:.3}",
8361                                    pos + k0 + k,
8362                                    target,
8363                                    lse - lg[target] as f64,
8364                                    top,
8365                                    lg[target],
8366                                    lg[top]
8367                                );
8368                            }
8369                        }
8370                        k0 = k1;
8371                    }
8372                    pos = end;
8373                }
8374                return Ok((nll, cnt));
8375            }
8376            for pos in 0..ids.len().saturating_sub(1) {
8377                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, task_mask);
8378                self.nll_check_graph("serial forward", pos)?;
8379                // Architectures whose head lives inside their own stack return
8380                // the logits out of band and a zero hidden — DeepSeek-V4 folds
8381                // its hyper-connection copies between the last layer and the
8382                // norm, so it cannot hand back a vector this loop could use.
8383                // Scoring the zeros gave a perplexity of exactly the vocabulary
8384                // size, which is a uniform distribution reported as a
8385                // measurement. `generate` already reads this channel.
8386                let out_of_band = self.graph_logits.take();
8387                if self.graph_head_required && out_of_band.is_none() {
8388                    METAL_GRAPH_HEAD_MISS.fetch_add(
8389                        1,
8390                        std::sync::atomic::Ordering::Relaxed,
8391                    );
8392                    return Err(format!(
8393                        "fused Metal graph head did not complete at NLL position {pos}"
8394                    ));
8395                }
8396                if pos < start {
8397                    continue;
8398                }
8399                let logits = match out_of_band {
8400                    Some(lg) => lg,
8401                    None => {
8402                        let normed = inference::rms_norm(
8403                            &hidden,
8404                            &self.weights.final_norm,
8405                            self.rms_eps,
8406                            self.norm_style,
8407                        );
8408                        // lm_head_forward applies the final-logit softcap itself
8409                        // — capping again here double-squashed gemma-class
8410                        // logits (tanh∘tanh) and reported a flattered ppl.
8411                        self.lm_head_forward(&normed)
8412                    }
8413                };
8414                let target = ids[pos + 1] as usize;
8415                let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8416                let lse: f64 = logits
8417                    .iter()
8418                    .map(|&v| ((v - max) as f64).exp())
8419                    .sum::<f64>()
8420                    .ln()
8421                    + max as f64;
8422                let tok_nll = lse - logits[target] as f64;
8423                if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8424                    let top = logits
8425                        .iter()
8426                        .enumerate()
8427                        .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8428                        .map(|(i, _)| i)
8429                        .unwrap_or(0);
8430                    eprintln!(
8431                        "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8432                        logits[target], logits[top]
8433                    );
8434                }
8435                nll += tok_nll;
8436                cnt += 1;
8437            }
8438            Ok((nll, cnt))
8439        })();
8440        self.nll_end();
8441        result
8442    }
8443
8444    /// Score one post-layer hidden with the same final norm/head path used by
8445    /// decode. Keeping this in one helper is important for the production
8446    /// batch scorer: its rows stop before the final norm, just like the
8447    /// per-position O(1) path below.
8448    fn nll_from_hidden(&mut self, hidden: &[f32], target: u32, pos: usize) -> f64 {
8449        let normed = inference::rms_norm(
8450            hidden,
8451            &self.weights.final_norm,
8452            self.rms_eps,
8453            self.norm_style,
8454        );
8455        // lm_head_forward applies the final-logit softcap itself — capping
8456        // again here double-squashed gemma-class logits in earlier scorers.
8457        let mut logits = self.lm_head_forward(&normed);
8458        let target = target as usize;
8459        let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8460        let lse: f64 = logits
8461            .iter()
8462            .map(|&v| ((v - max) as f64).exp())
8463            .sum::<f64>()
8464            .ln()
8465            + max as f64;
8466        let tok_nll = lse - logits[target] as f64;
8467        if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8468            let top = logits
8469                .iter()
8470                .enumerate()
8471                .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8472                .map(|(i, _)| i)
8473                .unwrap_or(0);
8474            eprintln!(
8475                "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8476                logits[target], logits[top]
8477            );
8478        }
8479        attention::recycle_buf(&mut logits);
8480        tok_nll
8481    }
8482
8483    /// `CMF_STATE_TRACE`: per-layer magnitude of the recurrent record at
8484    /// position `pos` — the whole `linear_state` (vmf: S then the conv
8485    /// ring; GDN: conv ring then S), its recurrent S part alone, the
8486    /// bounded ring, and the last per-position KV row count.  The tool
8487    /// that separated the ~4k perplexity cliff of the 500-step exports
8488    /// (a state that keeps climbing past the trained window) from a
8489    /// runtime boundary; one line per layer, `STATE pos=… layer=…`.
8490    fn trace_recurrent_state(&self, pos: usize) {
8491        let stats = |v: &[f32]| -> (f64, f64) {
8492            if v.is_empty() {
8493                return (0.0, 0.0);
8494            }
8495            let (mut ss, mut mx) = (0f64, 0f64);
8496            for &x in v {
8497                ss += (x as f64) * (x as f64);
8498                mx = mx.max((x as f64).abs());
8499            }
8500            ((ss / v.len() as f64).sqrt(), mx)
8501        };
8502        for (li, l) in self.kv_cache.layers.iter().enumerate() {
8503            let lw = &self.weights.layers[self.phys_layer(li)];
8504            let (kind, s_len) = match &lw.attn {
8505                AttnKind::Linear(_) => (
8506                    "vmf",
8507                    self.vmf_cfg.map(|c| c.state_len()).unwrap_or(0),
8508                ),
8509                AttnKind::LinearGdn(_) => (
8510                    "gdn",
8511                    self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0),
8512                ),
8513                AttnKind::Bounded(_) => ("bounded", 0),
8514                AttnKind::Full { .. } => ("full", 0),
8515                _ => ("other", 0),
8516            };
8517            let (rms, max) = stats(&l.linear_state);
8518            let s_part = if kind == "vmf" {
8519                &l.linear_state[..s_len.min(l.linear_state.len())]
8520            } else if kind == "gdn" {
8521                let ring = l.linear_state.len().saturating_sub(s_len).min(l.linear_state.len());
8522                &l.linear_state[ring..]
8523            } else {
8524                &l.linear_state[..0]
8525            };
8526            let (s_rms, s_max) = stats(s_part);
8527            let (ring_rms, ring_len) = match &l.bounded {
8528                Some(b) => (stats(&b.ring_k).0, b.ring_k.len()),
8529                None => (0.0, 0),
8530            };
8531            eprintln!(
8532                "STATE pos={pos} layer={li} kind={kind} state_len={} rms={rms:.5} max={max:.4} \
8533                 S_rms={s_rms:.5} S_max={s_max:.4} ring_k_rms={ring_rms:.5} ring_len={ring_len} kv_rows={}",
8534                l.linear_state.len(),
8535                l.seq_len
8536            );
8537        }
8538    }
8539
8540    /// Teacher-forced NLL of the CONVERTED model: the O(1) Nyström path
8541    /// is ACTIVE over the scored positions. Returns `Ok((nll sum, scored
8542    /// count))` over `prefill..len-1` and surfaces a post-mutation batch
8543    /// failure instead of returning a partial score.
8544    ///
8545    /// Runtime discipline, deliberately NOT the matrix probe's: the
8546    /// requested prefix plus any required deferred lead-in run the exact
8547    /// prompt pass — that pass is what freezes the landmarks and M — and
8548    /// every post-seal scored position goes through `NystromState::step()`,
8549    /// the same code decode runs.
8550    /// So the landmarks are PREFILL-frozen (what ships), not
8551    /// full-sequence oracles (what the published probe measured). When the
8552    /// requested prefix is shorter than the bounded transition, rows in the
8553    /// exact lead-in are still scored so the shifted target range is stable.
8554    ///
8555    /// Pair with `nll_ids_from(ids, prefill)` for the exact baseline
8556    /// over the identical token set — that ratio is the honest one.
8557    pub fn nll_ids_o1(&mut self, ids: &[u32], prefill: usize) -> Result<(f64, usize), String> {
8558        // This scorer consumes host hiddens, so never request the optional
8559        // token-graph lm_head side channel. `nll_begin` also consumes a
8560        // prior graph failure and clears only the cancel bit that failure
8561        // raised, leaving a caller-owned cancellation observable.
8562        self.nll_begin()?;
8563        let requested_prefix = (prefill > 0).then_some(prefill);
8564        self.o1_begin_with_prefix(requested_prefix);
8565        let n = ids.len().saturating_sub(1);
8566        let requested_start = prefill.min(n);
8567        // The exact prefix must reach the deferred boundary before a
8568        // collecting layer can convert. Rows between the requested start and
8569        // that boundary remain part of the public NLL range and are scored
8570        // from the same hidden pass below.
8571        let exact_end = if self.o1_active() {
8572            match requested_prefix {
8573                Some(requested) => self.o1_effective_boundary(requested),
8574                None => self
8575                    .o1_cfg
8576                    .as_ref()
8577                    .and_then(|c| crate::nystrom::o1_deferred_boundary(c.w, c.sink)),
8578            }
8579            .unwrap_or(requested_start)
8580            .min(n)
8581        } else {
8582            requested_start
8583        };
8584        let mut nll = 0f64;
8585        let mut cnt = 0usize;
8586
8587        // Exact prompt pass over ids[..exact_end]: the seal consumes its
8588        // q/k/v. Rows at or after requested_start are scored here when the
8589        // bounded lead-in is longer than the caller's requested prefix.
8590        let mut pos = 0usize;
8591        if self.can_prefill_batched() {
8592            const CHUNK: usize = 128;
8593            while pos < exact_end {
8594                let end = (pos + CHUNK).min(exact_end);
8595                let hiddens = self.prefill_batch(&ids[pos..end], pos);
8596                if self
8597                    .graph_failed
8598                    .swap(false, std::sync::atomic::Ordering::Relaxed)
8599                {
8600                    self.cancel
8601                        .store(false, std::sync::atomic::Ordering::Relaxed);
8602                    self.nll_end();
8603                    return Err("GPU graph failed during O(1) NLL prefix".into());
8604                }
8605                for row in 0..end - pos {
8606                    let score_pos = pos + row;
8607                    if score_pos >= requested_start && score_pos < n {
8608                        nll += self.nll_from_hidden(
8609                            &hiddens[row * self.hidden_size..(row + 1) * self.hidden_size],
8610                            ids[score_pos + 1],
8611                            score_pos,
8612                        );
8613                        cnt += 1;
8614                    }
8615                }
8616                pos = end;
8617            }
8618        } else {
8619            while pos < exact_end {
8620                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8621                if self
8622                    .graph_failed
8623                    .swap(false, std::sync::atomic::Ordering::Relaxed)
8624                {
8625                    self.cancel
8626                        .store(false, std::sync::atomic::Ordering::Relaxed);
8627                    self.nll_end();
8628                    return Err("GPU graph failed during O(1) NLL prefix".into());
8629                }
8630                if pos >= requested_start && pos < n {
8631                    nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8632                    cnt += 1;
8633                }
8634                pos += 1;
8635            }
8636        }
8637        self.o1_seal_checked().map_err(|err| {
8638            self.nll_end();
8639            err
8640        })?;
8641
8642        // Reuse the production whole-token batch graph for the post-seal
8643        // suffix when the caller explicitly enabled both routes. This is a
8644        // teacher-forced scorer, so every row is ids[pos] and its target is
8645        // ids[pos + 1]; no speculative tail or rollback state is involved.
8646        // A first Declined is safe to handle with the established serial O(1)
8647        // path. Once a chunk completes, however, the device recurrent state
8648        // owns the sequence and a later decline must be terminal rather than
8649        // falling back to stale CPU state.
8650        let batch_k = std::env::var("CMF_BATCH_K")
8651            .ok()
8652            .and_then(|v| v.parse::<usize>().ok())
8653            .unwrap_or(0);
8654        let batch_admitted = batch_k > 0
8655            && self.can_prefill_batched()
8656            && self.o1_active()
8657            && std::env::var("CMF_O1_GPU").as_deref() == Ok("1")
8658            && (0..self.num_layers).all(|li| {
8659                let cache = &self.kv_cache.layers[self.phys_layer(li)];
8660                cache.o1.is_none() || cache.o1_views().is_some()
8661            });
8662        if std::env::var("CMF_GRAPH_PROF").is_ok() {
8663            eprintln!(
8664                "nll-batch: phase=post-seal admission={} requested_k={} scored_rows={}",
8665                batch_admitted,
8666                batch_k,
8667                n.saturating_sub(exact_end),
8668            );
8669        }
8670        let mut batch_completed = false;
8671        if batch_admitted && exact_end < n {
8672            let hs = self.hidden_size;
8673            let mut batch_pos = exact_end;
8674            while batch_pos < n {
8675                let end = (batch_pos + batch_k).min(n);
8676                let bk = end - batch_pos;
8677                let mut hiddens = vec![0.0f32; bk * hs];
8678                for (row, &id) in ids[batch_pos..end].iter().enumerate() {
8679                    hiddens[row * hs..(row + 1) * hs].copy_from_slice(&self.embed_single(id));
8680                }
8681                let positions: Vec<usize> = (batch_pos..end).collect();
8682                let t_batch = std::time::Instant::now();
8683                let outcome = self.try_batch_graph_wgpu(&mut hiddens, &positions, bk, None);
8684                if std::env::var("CMF_GRAPH_PROF").is_ok() {
8685                    let ms = t_batch.elapsed().as_secs_f64() * 1000.0;
8686                    eprintln!(
8687                        "nll-batch: phase=post-seal mode=o1 k={bk} pos={}..{} outcome={outcome:?} {ms:.1} ms ({:.1} tok/s)",
8688                        batch_pos,
8689                        end.saturating_sub(1),
8690                        bk as f64 / (ms / 1000.0),
8691                    );
8692                }
8693                if let Err(err) = self.nll_check_graph("batch graph", batch_pos) {
8694                    self.nll_end();
8695                    return Err(err);
8696                }
8697                match outcome {
8698                    crate::gpu::BatchGraphOutcome::Completed => {
8699                        batch_completed = true;
8700                        for row in 0..bk {
8701                            nll += self.nll_from_hidden(
8702                                &hiddens[row * hs..(row + 1) * hs],
8703                                ids[batch_pos + row + 1],
8704                                batch_pos + row,
8705                            );
8706                            cnt += 1;
8707                        }
8708                        batch_pos = end;
8709                    }
8710                    crate::gpu::BatchGraphOutcome::Declined => {
8711                        if batch_completed {
8712                            self.nll_end();
8713                            return Err(format!(
8714                                "O(1) NLL batch declined after completed chunk at position {batch_pos}"
8715                            ));
8716                        }
8717                        break;
8718                    }
8719                    crate::gpu::BatchGraphOutcome::Failed => {
8720                        self.nll_end();
8721                        return Err(format!(
8722                            "O(1) NLL batch graph failed after admission at position {batch_pos}"
8723                        ));
8724                    }
8725                }
8726            }
8727            if batch_completed && cnt == n.saturating_sub(requested_start) {
8728                self.nll_end();
8729                return Ok((nll, cnt));
8730            }
8731        }
8732
8733        // Serial O(1) fallback/reference. It is intentionally retained when
8734        // batch admission declines before mutation; callers must label this
8735        // CMF_BATCH_K=0/per-position path separately from the production
8736        // whole-token batch route.
8737        for pos in exact_end..n {
8738            let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8739            if self
8740                .graph_failed
8741                .swap(false, std::sync::atomic::Ordering::Relaxed)
8742            {
8743                self.cancel
8744                    .store(false, std::sync::atomic::Ordering::Relaxed);
8745                self.nll_end();
8746                return Err(format!(
8747                    "GPU graph failed during O(1) NLL serial scoring at position {pos}"
8748                ));
8749            }
8750            nll += self.nll_from_hidden(&hidden, ids[pos + 1], pos);
8751            cnt += 1;
8752        }
8753        self.nll_end();
8754        Ok((nll, cnt))
8755    }
8756
8757    /// Teacher-forced calibration data (B1): for each position, whether the
8758    /// argmax equals the actual next token, and the top-1 softmax prob
8759    /// (top-1 probability) under EACH temperature in `temps` — all from ONE forward
8760    /// pass (argmax/correctness are temperature-invariant; only p_max
8761    /// reshapes). Feeds `cortiq calibrate` (reliability/ECE + temperature
8762    /// fit): is the model's confidence a true property, or does it need a
8763    /// measured scaling?
8764    pub fn calib_ids(&mut self, ids: &[u32], temps: &[f32]) -> (Vec<bool>, Vec<Vec<f32>>) {
8765        self.clear_sequence_state();
8766        let n = ids.len().saturating_sub(1);
8767        let mut correct = Vec::with_capacity(n);
8768        let mut pmax = Vec::with_capacity(n);
8769        for pos in 0..n {
8770            let emb = self.embed_single(ids[pos]);
8771            let hidden = self.forward_layers(&emb, pos, None);
8772            let logits = if let Some(logits) = self.graph_logits.take() {
8773                logits
8774            } else {
8775                let normed = inference::rms_norm(
8776                    &hidden,
8777                    &self.weights.final_norm,
8778                    self.rms_eps,
8779                    self.norm_style,
8780                );
8781                // lm_head_forward applies the final-logit softcap itself —
8782                // capping again here double-squashed gemma-class logits
8783                // (tanh∘tanh) and reported a flattered ppl.
8784                self.lm_head_forward(&normed)
8785            };
8786            let target = ids[pos + 1] as usize;
8787            let (mut amax, mut mval) = (0usize, f32::NEG_INFINITY);
8788            for (i, &v) in logits.iter().enumerate() {
8789                if v > mval {
8790                    mval = v;
8791                    amax = i;
8792                }
8793            }
8794            correct.push(amax == target);
8795            let row: Vec<f32> = temps
8796                .iter()
8797                .map(|&t| {
8798                    let tt = t.max(1e-3);
8799                    let s: f32 = logits.iter().map(|&v| ((v - mval) / tt).exp()).sum();
8800                    1.0 / s.max(1e-12) // numerator at the max is exp(0)=1
8801                })
8802                .collect();
8803            pmax.push(row);
8804        }
8805        self.clear_sequence_state();
8806        (correct, pmax)
8807    }
8808
8809    /// Teacher-forced PPL with the dynamic router driving per-window
8810    /// skill switches (VMF experiment №2 measurement). Sequential (φ
8811    /// must update per token), returns (ppl, switch_count). The router
8812    /// must be enabled (`enable_dynamic_routing`); else this equals
8813    /// plain `ppl_ids`. The active skill when scoring token t shapes the
8814    /// logits for t+1 — on-policy over the held-out text itself.
8815    pub fn ppl_ids_dynamic(&mut self, ids: &[u32]) -> Result<(f64, usize), String> {
8816        if self.dyn_router.is_none() {
8817            return Ok((self.ppl_ids(ids)?, 0));
8818        }
8819        self.nll_begin()?;
8820        let saved_active = self.dyn_active;
8821        let mut router = self
8822            .dyn_router
8823            .take()
8824            .ok_or_else(|| "dynamic router disappeared before PPL scoring".to_string())?;
8825        router.reset();
8826        self.dyn_phi_seen = 0;
8827        let _ = self.set_active_skill(None);
8828
8829        let result: Result<(f64, usize), String> = (|| {
8830            let mut nll = 0f64;
8831            let mut cnt = 0usize;
8832            for pos in 0..ids.len().saturating_sub(1) {
8833                let hidden = self.forward_layers(&self.embed_single(ids[pos]), pos, None);
8834                self.nll_check_graph("dynamic serial forward", pos)?;
8835                let out_of_band = self.graph_logits.take();
8836                let mut logits = match out_of_band {
8837                    Some(lg) => lg,
8838                    None => {
8839                        let normed = inference::rms_norm(
8840                            &hidden,
8841                            &self.weights.final_norm,
8842                            self.rms_eps,
8843                            self.norm_style,
8844                        );
8845                        // lm_head_forward applies the final-logit softcap itself —
8846                        // capping again here double-squashed gemma-class logits
8847                        // and reported a flattered ppl.
8848                        self.lm_head_forward(&normed)
8849                    }
8850                };
8851                let target = ids[pos + 1] as usize;
8852                let max = logits.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
8853                let lse: f64 = logits
8854                    .iter()
8855                    .map(|&v| ((v - max) as f64).exp())
8856                    .sum::<f64>()
8857                    .ln()
8858                    + max as f64;
8859                let tok_nll = lse - logits[target] as f64;
8860                if std::env::var("CMF_PPL_TRACE").is_ok() && pos < 48 {
8861                    let top = logits
8862                        .iter()
8863                        .enumerate()
8864                        .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
8865                        .map(|(i, _)| i)
8866                        .unwrap_or(0);
8867                    eprintln!(
8868                        "pos {pos:3} tgt {target:6} nll {tok_nll:7.3} | top1 {top:6} lg[t]={:.2} lg[top]={:.2}",
8869                        logits[target], logits[top]
8870                    );
8871                }
8872                nll += tok_nll;
8873                cnt += 1;
8874                attention::recycle_buf(&mut logits);
8875                // Route on the evolving phi (drives the NEXT token's skill).
8876                let phi = self.dyn_phi_ema.clone();
8877                if let Some(new_active) = router.step(&phi, pos) {
8878                    let _ = self.set_active_skill(new_active);
8879                }
8880            }
8881            Ok(((nll / cnt.max(1) as f64).exp(), router.switches.len()))
8882        })();
8883
8884        // Restore the detached router and the active overlay on both success
8885        // and failure. The scoring state is cleared independently below.
8886        let _ = self.set_active_skill(saved_active);
8887        self.dyn_router = Some(router);
8888        self.nll_end();
8889        result
8890    }
8891
8892    /// Routing probe φ (spec §9): mean-pooled hidden after `layer`.
8893    pub fn probe_phi(&mut self, ids: &[u32], layer: usize) -> Vec<f32> {
8894        self.clear_sequence_state();
8895        let mut acc = vec![0f32; self.hidden_size];
8896        for (pos, &id) in ids.iter().enumerate() {
8897            let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8898            for (a, v) in acc.iter_mut().zip(&h) {
8899                *a += v;
8900            }
8901        }
8902        let n = ids.len().max(1) as f32;
8903        for a in acc.iter_mut() {
8904            *a /= n;
8905        }
8906        self.clear_sequence_state();
8907        acc
8908    }
8909
8910    /// Router-v2 φ probe (spec §9.4, `phi.pool = "span_mean"`): the hidden
8911    /// AFTER `layer` — the same per-position walk and the same quantity as
8912    /// [`Self::probe_phi`] — averaged over the positions in `span` only
8913    /// (the user text between the template's prefix and suffix ids), NOT
8914    /// unit-normalized (the decision normalizes). The walk stops at
8915    /// `span.end`: causality makes the later positions irrelevant, so the
8916    /// result is bit-identical to probing `ids[..span.end]`. An empty span
8917    /// gives the zero vector, which the decision treats as degenerate.
8918    ///
8919    /// Every sequence state is reset before and after — the host KV/ring/
8920    /// recurrent state, the reuse keys (`kv_history`, `kv_prefix`) and the
8921    /// device graph's sequence — so run it on a pipeline that does not
8922    /// also serve a conversation (its prefix reuse would be lost).
8923    pub fn probe_phi_span(
8924        &mut self,
8925        ids: &[u32],
8926        layer: usize,
8927        span: std::ops::Range<usize>,
8928    ) -> Vec<f32> {
8929        #[cfg(target_os = "macos")]
8930        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8931        let end = span.end.min(ids.len());
8932        let start = span.start.min(end);
8933        let reset = |p: &mut Self| p.clear_sequence_state();
8934        reset(self);
8935        let mut acc = vec![0f32; self.hidden_size];
8936        for (pos, &id) in ids[..end].iter().enumerate() {
8937            let h = self.forward_layers_upto(&self.embed_single(id), pos, None, Some(layer));
8938            if pos >= start {
8939                for (a, v) in acc.iter_mut().zip(&h) {
8940                    *a += v;
8941                }
8942            }
8943        }
8944        let n = end - start;
8945        if n > 0 {
8946            let n = n as f32;
8947            for a in acc.iter_mut() {
8948                *a /= n;
8949            }
8950        }
8951        reset(self);
8952        acc
8953    }
8954
8955    /// One decode step of the current sequence: forward `token` at
8956    /// `position` (the cache holds positions `[0, position)`, e.g. after
8957    /// [`Self::forward_ids`]) and return the next-token logits — the same
8958    /// forward and head the generation loop runs (resident-graph logits
8959    /// when the graph ran, final norm + lm_head otherwise). The logit-dump
8960    /// tools drive greedy decoding with it so every position is observable.
8961    pub fn decode_step_logits(&mut self, token: u32, position: usize) -> Vec<f32> {
8962        #[cfg(target_os = "macos")]
8963        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
8964        self.graph_logits = None;
8965        let hidden = self.forward_layers(&self.embed_single(token), position, None);
8966        if let Some(logits) = self.graph_logits.take() {
8967            return logits;
8968        }
8969        inference::rms_norm_into(
8970            &hidden,
8971            &self.weights.final_norm,
8972            self.rms_eps,
8973            self.norm_style,
8974            &mut self.ws.n1,
8975        );
8976        self.lm_head_forward(&self.ws.n1)
8977    }
8978
8979    /// Layer-major batched prefill (prefill-GEMM): full-attention —
8980    /// per-position with the existing operators (KV grows naturally,
8981    /// causality preserved), GDN projections / FFN / MoE — batched
8982    /// (a weight row is read from DRAM once per chunk, not per
8983    /// position). Returns the hidden of all positions [b × hidden].
8984    fn prefill_batch(&mut self, ids: &[u32], start_pos: usize) -> Vec<f32> {
8985        self.prefill_batch_masked(ids, start_pos, None)
8986    }
8987
8988    /// `prefill_batch` with a task mask honored on the dense-FFN panels
8989    /// (the masked-inference fast path: full fused compute, mask lands on
8990    /// the activations). The whole-chunk GPU graph is skipped for masked
8991    /// layers by the callers' arms; the per-GEMM device paths stay in
8992    /// play because the zeroing happens on the host between them.
8993    fn prefill_batch_masked(
8994        &mut self,
8995        ids: &[u32],
8996        start_pos: usize,
8997        task_mask: Option<&TaskMask>,
8998    ) -> Vec<f32> {
8999        self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, usize::MAX)
9000    }
9001
9002    /// One prompt chunk through the whole stack, post-stack rows out (no
9003    /// final norm) — the ingest generation uses, shared by scoring and
9004    /// `forward_ids` so they measure the same execution: the batched wgpu
9005    /// graph's device prefix plus the host's batched walk for the rest when
9006    /// `batch_prefix_prefill` holds and the graph admits the chunk, else
9007    /// the host's chunked prefill. Err only when a graph that had mutated
9008    /// device state failed.
9009    fn prefill_rows(
9010        &mut self,
9011        ids: &[u32],
9012        pos: usize,
9013        task_mask: Option<&TaskMask>,
9014    ) -> Result<Vec<f32>, String> {
9015        self.prefill_input_rows(PrefillIn::Ids(ids), pos, task_mask)
9016    }
9017
9018    fn prefill_input_rows(
9019        &mut self,
9020        input: PrefillIn<'_>,
9021        pos: usize,
9022        task_mask: Option<&TaskMask>,
9023    ) -> Result<Vec<f32>, String> {
9024        self.mimo_moe_prepare();
9025        let hs = self.hidden_size;
9026        let bk = match input {
9027            PrefillIn::Ids(ids) => ids.len(),
9028            PrefillIn::Hidden(rows) => rows.len() / hs,
9029        };
9030        #[cfg(not(target_os = "macos"))]
9031        if task_mask.is_none()
9032            && !self.o1_active()
9033            && bk > 1
9034            && (self.batch_prefix_prefill()
9035                || (self.verify_exact_moe
9036                    && crate::gpu::enabled_here()
9037                    && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)))
9038        {
9039            let mut hiddens = match input {
9040                PrefillIn::Hidden(rows) => rows.to_vec(),
9041                PrefillIn::Ids(ids) => ids.iter().flat_map(|&id| self.embed_single(id)).collect(),
9042            };
9043            let positions: Vec<usize> = (pos..pos + bk).collect();
9044            let mut run = 0usize;
9045            match self.try_batch_graph_wgpu_prefix(
9046                &mut hiddens,
9047                &positions,
9048                bk,
9049                None,
9050                Some(&mut run),
9051            ) {
9052                crate::gpu::BatchGraphOutcome::Completed => {
9053                    let out = if run < self.num_layers {
9054                        self.prefill_batch_span(
9055                            PrefillIn::Hidden(&hiddens),
9056                            pos,
9057                            None,
9058                            run,
9059                            self.num_layers,
9060                        )
9061                    } else {
9062                        hiddens
9063                    };
9064                    return if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
9065                        Err("MiMo attention graph failed after admission".into())
9066                    } else { Ok(out) };
9067                }
9068                crate::gpu::BatchGraphOutcome::Failed => {
9069                    return Err("batched prefix prefill failed after admission".into());
9070                }
9071                crate::gpu::BatchGraphOutcome::Declined => {
9072                    // Rows an earlier chunk left on the device only.
9073                    #[cfg(feature = "gpu")]
9074                    self.pull_lagging_host_kv(0, self.num_layers, pos);
9075                }
9076            }
9077        }
9078        let out = self.prefill_batch_span(input, pos, task_mask, 0, usize::MAX);
9079        if self.graph_failed.load(std::sync::atomic::Ordering::Relaxed) {
9080            Err("batch tail graph failed after admission".into())
9081        } else { Ok(out) }
9082    }
9083
9084    /// The layer-major batched walk over a layer span [from..upto_excl):
9085    /// the whole prefill machinery (chunk graph, batched attends, GEMM
9086    /// panels) for a PARTIAL stack — the network split's prefill rides
9087    /// the same canon as the local one. Input is token ids (embeds
9088    /// itself, coordinator side) or ready boundary hiddens (worker side).
9089    fn prefill_batch_span(
9090        &mut self,
9091        input: PrefillIn<'_>,
9092        start_pos: usize,
9093        task_mask: Option<&TaskMask>,
9094        from: usize,
9095        upto_excl: usize,
9096    ) -> Vec<f32> {
9097        let hs = self.hidden_size;
9098        let b = match input {
9099            PrefillIn::Ids(ids) => ids.len(),
9100            PrefillIn::Hidden(hb) => hb.len() / hs,
9101        };
9102        let upto_excl = upto_excl.min(self.num_layers);
9103        // The CPU embed is deferred: when the chunk graph takes the run
9104        // from layer 0 it gathers the embeddings on the device instead.
9105        // A hidden input is ready by definition.
9106        let mut h: Vec<f32>;
9107        let mut h_ready;
9108        match input {
9109            PrefillIn::Ids(_) => {
9110                h = vec![0.0; b * hs];
9111                h_ready = false;
9112            }
9113            PrefillIn::Hidden(hb) => {
9114                h = hb.to_vec();
9115                h_ready = true;
9116            }
9117        }
9118        let fill_h = |h: &mut Vec<f32>, me: &Self| {
9119            if let PrefillIn::Ids(ids) = input {
9120                for (bi, &id) in ids.iter().enumerate() {
9121                    let e = me.embed_single(id);
9122                    h[bi * hs..(bi + 1) * hs].copy_from_slice(&e);
9123                }
9124                if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9125                    if let Ok(t) = tp.parse::<usize>() {
9126                        if t >= start_pos && t < start_pos + ids.len() {
9127                            let bi = t - start_pos;
9128                            let row = &h[bi * hs..(bi + 1) * hs];
9129                            let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9130                            eprintln!(
9131                                "BATCH pos {t} embed: id {} |h| = {n:.6} h0 {:.6} h1 {:.6} | b={} start={start_pos} ids[..8]={:?}",
9132                                ids[bi],
9133                                row[0],
9134                                row[1],
9135                                ids.len(),
9136                                &ids[..ids.len().min(8)]
9137                            );
9138                        }
9139                    }
9140                }
9141            }
9142        };
9143        let (_nkv, _hd, _rd, eps) = (
9144            self.num_kv_heads,
9145            self.head_dim,
9146            self.rotary_dim,
9147            self.rms_eps,
9148        );
9149        let pool = self.pool.clone();
9150        let norm_style = self.norm_style;
9151        self.mimo_moe_prepare();
9152        let automatic_gpu_prefix = self.automatic_gpu_prefix();
9153
9154        #[cfg(target_os = "macos")]
9155        let mut chunk_skip_until = 0usize;
9156        for li in from..upto_excl {
9157            let _capacity_tail = automatic_gpu_prefix
9158                .filter(|&prefix| {
9159                    li >= prefix && !(self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false))
9160                })
9161                .map(|_| crate::gpu::enter_cpu_scope());
9162            crate::gpu::set_layer(li as i64); // layer-split GPU/CPU
9163            // Spark-X2.5: the layer's int8 matrix-unit GEMMs take the
9164            // `zi_mm` tile, bit-identical to the 64x64 kernel and about
9165            // three times its rate (`gpu_wgpu::zi_gemm_route`).
9166            #[cfg(feature = "gpu")]
9167            let _fast_gemm = self
9168                .proj_gate_sigmoid
9169                .then(crate::gpu::enter_prefill_fast_gemm);
9170            // GPU chunk graph (default-on under CMF_GPU=1): a run of
9171            // consecutive eligible layers for the whole chunk in ONE
9172            // Metal submission — norm, QKV, RoPE with fused mirror
9173            // append, causal attend, O, FFN, hidden device-resident
9174            // across the run. Any refusal falls through to the CPU path.
9175            #[cfg(target_os = "macos")]
9176            if task_mask.is_none() {
9177                if li < chunk_skip_until {
9178                    continue;
9179                }
9180                // Device-side embedding needs a q8_row embedding matrix;
9181                // with any other layout the CPU fills `h` first and the
9182                // graph starts from a ready hidden (refusing the whole
9183                // run over the embedding alone kept q4t models — the
9184                // whole Nanbeige/Bonsai class — on the CPU prefill).
9185                if !h_ready && li == 0 && self.weights.embed_tokens.q8_row_parts().is_none() {
9186                    fill_h(&mut h, self);
9187                    h_ready = true;
9188                }
9189                let ids_for_embed = match input {
9190                    PrefillIn::Ids(ids) => (!h_ready && li == 0).then_some(ids),
9191                    PrefillIn::Hidden(_) => None,
9192                };
9193                let end = self.chunk_run_gpu(li, &mut h, b, start_pos, ids_for_embed, upto_excl);
9194                if end > li {
9195                    h_ready = true;
9196                    chunk_skip_until = end;
9197                    // Looped Transformer: the graph stopped at a loop
9198                    // boundary — apply final norm before the next iteration.
9199                    if self.is_loop_end(end - 1) && end < self.num_layers {
9200                        for bi in 0..b {
9201                            let normed = inference::rms_norm(
9202                                &h[bi * hs..(bi + 1) * hs],
9203                                &self.weights.final_norm,
9204                                eps,
9205                                norm_style,
9206                            );
9207                            h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9208                        }
9209                    }
9210                    continue;
9211                }
9212            }
9213            if !h_ready {
9214                fill_h(&mut h, self);
9215                h_ready = true;
9216            }
9217            if task_mask.is_none() && self.verify_exact_moe {
9218                let positions: Vec<_> = (start_pos..start_pos + b).collect();
9219                match self.mimo_graph_layer_rows(li, &mut h, &positions) {
9220                    crate::gpu::BatchGraphOutcome::Completed => continue,
9221                    crate::gpu::BatchGraphOutcome::Failed => return h,
9222                    crate::gpu::BatchGraphOutcome::Declined => {},
9223                }
9224            }
9225            #[cfg(feature = "gpu")]
9226            self.pull_lagging_host_kv(li, li + 1, start_pos);
9227            let lw = &self.weights.layers[self.phys_layer(li)];
9228            let t_attn = std::time::Instant::now();
9229            // ── attention ──
9230            match &lw.attn {
9231                AttnKind::Kda(w) => {
9232                    // Projections batched, recurrence sequential.
9233                    let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
9234                    let mut normed = vec![0.0f32; b * hs];
9235                    for bi in 0..b {
9236                        inference::rms_norm_into(
9237                            &h[bi * hs..(bi + 1) * hs],
9238                            &lw.input_norm,
9239                            eps,
9240                            norm_style,
9241                            &mut normed[bi * hs..(bi + 1) * hs],
9242                        );
9243                    }
9244                    let attn = crate::linear_core::kda_forward_batch(
9245                        &normed,
9246                        b,
9247                        w,
9248                        &cfg,
9249                        &mut self.kv_cache.layers[li].linear_state,
9250                        pool.as_deref(),
9251                    );
9252                    for (dst, &a) in h.iter_mut().zip(&attn) {
9253                        *dst += a;
9254                    }
9255                }
9256                AttnKind::LinearGdn(w) => {
9257                    // Projections batched, recurrence sequential.
9258                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
9259                    let mut normed = vec![0.0f32; b * hs];
9260                    for bi in 0..b {
9261                        let r = inference::rms_norm(
9262                            &h[bi * hs..(bi + 1) * hs],
9263                            &lw.input_norm,
9264                            eps,
9265                            norm_style,
9266                        );
9267                        normed[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9268                    }
9269                    let attn = crate::linear_core::gdn_forward_batch(
9270                        &normed,
9271                        b,
9272                        w,
9273                        &cfg,
9274                        &mut self.kv_cache.layers[li].linear_state,
9275                        pool.as_deref(),
9276                    );
9277                    for (dst, &a) in h.iter_mut().zip(&attn) {
9278                        *dst += a;
9279                    }
9280                }
9281                AttnKind::ShortConv(w) => {
9282                    // Projections batched over the chunk; the conv walks the
9283                    // contiguous positions in order (same ring as decode).
9284                    let cfg = self
9285                        .short_conv_cfg
9286                        .expect("short-conv layer without short_conv_cfg");
9287                    let mut normed = vec![0.0f32; b * hs];
9288                    for bi in 0..b {
9289                        inference::rms_norm_into(
9290                            &h[bi * hs..(bi + 1) * hs],
9291                            &lw.input_norm,
9292                            eps,
9293                            norm_style,
9294                            &mut normed[bi * hs..(bi + 1) * hs],
9295                        );
9296                    }
9297                    let attn = short_conv_forward_batch(
9298                        &normed,
9299                        b,
9300                        w,
9301                        &cfg,
9302                        &mut self.kv_cache.layers[li].linear_state,
9303                        pool.as_deref(),
9304                    );
9305                    for (dst, &a) in h.iter_mut().zip(&attn) {
9306                        *dst += a;
9307                    }
9308                }
9309                AttnKind::Mla(w) => {
9310                    // Per-position prefill (correctness first; latent
9311                    // batching is a later optimization).
9312                    let inv_freq_l = self.layer_inv_freq(li);
9313                    let rs = self.layer_rope_scale(li);
9314                    let mut normed = vec![0.0f32; hs];
9315                    for bi in 0..b {
9316                        inference::rms_norm_into(
9317                            &h[bi * hs..(bi + 1) * hs],
9318                            &lw.input_norm,
9319                            eps,
9320                            norm_style,
9321                            &mut normed,
9322                        );
9323                        let ao = mla_attention(
9324                            w,
9325                            &normed,
9326                            &mut self.kv_cache.layers[li],
9327                            start_pos + bi,
9328                            &inv_freq_l,
9329                            rs,
9330                            eps,
9331                            pool.as_deref(),
9332                        );
9333                        for (dst, &a) in h[bi * hs..(bi + 1) * hs].iter_mut().zip(&ao) {
9334                            *dst += a;
9335                        }
9336                    }
9337                }
9338                AttnKind::Full {
9339                    wq,
9340                    wk,
9341                    wv,
9342                    wo,
9343                    q_norm,
9344                    k_norm,
9345                    output_gate,
9346                    softplus_gate,
9347                    bias,
9348                } => {
9349                    // Chunk-GEMM QKV/O; per-position causal attention
9350                    // inside (roadmap §3 P0 — full-attention prefill no
9351                    // longer re-reads the projection weights b times).
9352                    let mut normed = vec![0.0f32; b * hs];
9353                    for bi in 0..b {
9354                        inference::rms_norm_into(
9355                            &h[bi * hs..(bi + 1) * hs],
9356                            &lw.input_norm,
9357                            eps,
9358                            norm_style,
9359                            &mut normed[bi * hs..(bi + 1) * hs],
9360                        );
9361                    }
9362                    let inv_freq_l = self.layer_inv_freq(li);
9363                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
9364                    let cfg = QwenAttnCfg {
9365                        num_heads: self.layer_num_heads(li),
9366                        num_kv_heads: nkv_l,
9367                        head_dim: hd_l,
9368                        hidden_size: hs,
9369                        position: start_pos,
9370                        inv_freq: &inv_freq_l,
9371                        rotary_dim: rd_l,
9372                        scale: self.attn_scale,
9373                        softcap: self.attn_softcap,
9374                        window: self.layer_window(li),
9375                        v_norm: self.attn_v_norm,
9376                        qk_norm_after_rope: self.qk_norm_after_rope,
9377                        gate_sigmoid: self.proj_gate_sigmoid,
9378                        q_norm: q_norm.as_deref(),
9379                        k_norm: k_norm.as_deref(),
9380                        output_gate: *output_gate,
9381                        softplus_gate: softplus_gate
9382                            .as_ref()
9383                            .map(|(gate, per_head)| (gate, *per_head)),
9384                        rope_scale: self.layer_rope_scale(li),
9385                        bias: bias
9386                            .as_ref()
9387                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
9388                        rms_eps: eps,
9389                        norm_style,
9390                        pool: pool.as_deref(),
9391                        v_head_dim: self.layer_v_dim(li),
9392                    };
9393                    #[cfg(feature = "gpu")]
9394                    let _mirror = self
9395                        .prefill_mirror_target(li)
9396                        .map(crate::gpu::enter_prefill_mirror);
9397                    let mut attn = attention::qwen_attention_batch(
9398                        &normed,
9399                        b,
9400                        wq,
9401                        wk,
9402                        wv,
9403                        wo,
9404                        &mut self.kv_cache.layers[li],
9405                        &cfg,
9406                    );
9407                    if let Some(w) = &lw.attn_out_norm {
9408                        for bi in 0..b {
9409                            inference::rms_norm_into(
9410                                &attn[bi * hs..(bi + 1) * hs],
9411                                w,
9412                                eps,
9413                                norm_style,
9414                                &mut normed[bi * hs..(bi + 1) * hs],
9415                            );
9416                        }
9417                        attn.copy_from_slice(&normed);
9418                    }
9419                    for (dst, &a) in h.iter_mut().zip(&attn) {
9420                        *dst += a;
9421                    }
9422                }
9423                AttnKind::Bounded(w) => {
9424                    // Chunk-GEMM projections, the bounded operator per
9425                    // position over ring + chunk — never a growing KV.
9426                    let mut normed = vec![0.0f32; b * hs];
9427                    for bi in 0..b {
9428                        inference::rms_norm_into(
9429                            &h[bi * hs..(bi + 1) * hs],
9430                            &lw.input_norm,
9431                            eps,
9432                            norm_style,
9433                            &mut normed[bi * hs..(bi + 1) * hs],
9434                        );
9435                    }
9436                    let rope = self
9437                        .bounded_rope
9438                        .clone()
9439                        .expect("bounded layer without an installed rotation table");
9440                    let cfg = crate::bounded::BoundedAttnCfg {
9441                        num_heads: self.num_heads,
9442                        num_kv_heads: self.num_kv_heads,
9443                        head_dim: self.head_dim,
9444                        hidden_size: hs,
9445                        scale: self.attn_scale,
9446                        rope: &rope,
9447                        pool: pool.as_deref(),
9448                    };
9449                    let mut attn = crate::bounded::bounded_attention_batch(
9450                        &normed,
9451                        b,
9452                        w,
9453                        &mut self.kv_cache.layers[li],
9454                        &cfg,
9455                    );
9456                    if let Some(wn) = &lw.attn_out_norm {
9457                        for bi in 0..b {
9458                            inference::rms_norm_into(
9459                                &attn[bi * hs..(bi + 1) * hs],
9460                                wn,
9461                                eps,
9462                                norm_style,
9463                                &mut normed[bi * hs..(bi + 1) * hs],
9464                            );
9465                        }
9466                        attn.copy_from_slice(&normed);
9467                    }
9468                    for (dst, &a) in h.iter_mut().zip(&attn) {
9469                        *dst += a;
9470                    }
9471                    attention::recycle_buf(&mut attn);
9472                }
9473                AttnKind::Linear(w) => {
9474                    for bi in 0..b {
9475                        let normed = inference::rms_norm(
9476                            &h[bi * hs..(bi + 1) * hs],
9477                            &lw.input_norm,
9478                            eps,
9479                            norm_style,
9480                        );
9481                        vmf_phase_forward(
9482                            &normed,
9483                            w,
9484                            &self.vmf_cfg.expect("linear layer without vmf_cfg"),
9485                            &mut self.kv_cache.layers[li].linear_state,
9486                            pool.as_deref(),
9487                        )
9488                        .iter()
9489                        .enumerate()
9490                        .for_each(|(i, &a)| h[bi * hs + i] += a);
9491                    }
9492                }
9493            }
9494
9495            // ── FFN batched ──
9496            let lw = &self.weights.layers[self.phys_layer(li)];
9497            let mut post = vec![0.0f32; b * hs];
9498            for bi in 0..b {
9499                let r =
9500                    inference::rms_norm(&h[bi * hs..(bi + 1) * hs], &lw.post_norm, eps, norm_style);
9501                post[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9502            }
9503            // A restrictive per-visit FFN row lands on the activations
9504            // inside the dense arm; an all-open row costs nothing.
9505            let mask_row = task_mask
9506                .filter(|m| m.ffn_active_count(li) < self.intermediate_size)
9507                .and_then(|m| m.ffn_masks.get(li))
9508                .map(|v| v.as_slice());
9509            let attn_ns = t_attn.elapsed().as_nanos() as u64;
9510            let t_ffn = std::time::Instant::now();
9511            let mut ffn = match &lw.ffn {
9512                FfnKind::Dense(d) if !d.segs.is_empty() => {
9513                    tube_ffn(d, &post, b, pool.as_deref(), mask_row)
9514                }
9515                FfnKind::Dense(d) => dense_ffn_batch(d, &post, b, pool.as_deref(), mask_row),
9516                FfnKind::Moe(m) if self.verify_exact_moe && self.mimo_moe.is_dynamic(li, false) => {
9517                    moe_ffn_banked_rows(&mut self.mimo_moe, li, m, &post, b, hs, pool.as_deref())
9518                }
9519                FfnKind::Moe(m) if self.verify_exact_moe => {
9520                    moe_ffn_rows_exact(m, &post, b, hs, pool.as_deref())
9521                }
9522                // Keep prompt expert panels off the projection arena and
9523                // use their routes to prime the model-wide bank.
9524                FfnKind::Moe(m) if self.mimo_moe.is_dynamic(li, false) => {
9525                    let before = m.stats.borrow().clone();
9526                    let out = crate::gpu::cpu_scope(|| {
9527                        moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None)
9528                    });
9529                    self.mimo_moe.prime(li, m, &before);
9530                    out
9531                }
9532                FfnKind::Moe(m) => moe_ffn_batch(m, &post, b, hs, pool.as_deref(), None),
9533                // Dual-branch layers run per position (the expert branch
9534                // reads the raw residual — nothing to batch yet).
9535                FfnKind::DenseMoe(dm) => {
9536                    let mut out = vec![0.0f32; b * hs];
9537                    for bi in 0..b {
9538                        let r = dense_moe_ffn(
9539                            dm,
9540                            &post[bi * hs..(bi + 1) * hs],
9541                            &h[bi * hs..(bi + 1) * hs],
9542                            eps,
9543                            norm_style,
9544                            pool.as_deref(),
9545                        );
9546                        out[bi * hs..(bi + 1) * hs].copy_from_slice(&r);
9547                    }
9548                    out
9549                }
9550            };
9551            if prefill_prof_on() {
9552                PREFILL_SPLIT[0].fetch_add(attn_ns, std::sync::atomic::Ordering::Relaxed);
9553                PREFILL_SPLIT[1].fetch_add(t_ffn.elapsed().as_nanos() as u64, std::sync::atomic::Ordering::Relaxed);
9554            }
9555            if let Some(w) = &lw.ffn_out_norm {
9556                for bi in 0..b {
9557                    inference::rms_norm_into(
9558                        &ffn[bi * hs..(bi + 1) * hs],
9559                        w,
9560                        eps,
9561                        norm_style,
9562                        &mut post[bi * hs..(bi + 1) * hs],
9563                    );
9564                }
9565                ffn.copy_from_slice(&post);
9566            }
9567            for (dst, &f) in h.iter_mut().zip(&ffn) {
9568                *dst += f;
9569            }
9570            if let Some(sc) = lw.layer_scale {
9571                for v in h.iter_mut() {
9572                    *v *= sc;
9573                }
9574            }
9575            // CMF_LAYER_DUMP: every position's hidden after layer li.
9576            if self.layer_dump.is_some() {
9577                for bi in 0..b {
9578                    self.dump_layer_row(start_pos + bi, li, &h[bi * hs..(bi + 1) * hs]);
9579                }
9580            }
9581            if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
9582                if let Ok(t) = tp.parse::<usize>() {
9583                    if t >= start_pos && t < start_pos + b {
9584                        let bi = t - start_pos;
9585                        let row = &h[bi * hs..(bi + 1) * hs];
9586                        let n: f32 = row.iter().map(|x| x * x).sum::<f32>().sqrt();
9587                        eprintln!(
9588                            "BATCH pos {t} after layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
9589                            row[0], row[1]
9590                        );
9591                    }
9592                }
9593            }
9594            // CMF_DEBUG_LAYERS=1: per-layer hidden-state health of the
9595            // LAST prompt position — the knife for "which layer type
9596            // breaks first" on a new architecture.
9597            if std::env::var("CMF_DEBUG_LAYERS").is_ok() {
9598                let row = &h[(b - 1) * hs..b * hs];
9599                let rms =
9600                    (row.iter().map(|&v| (v as f64) * (v as f64)).sum::<f64>() / hs as f64).sqrt();
9601                let mx = row.iter().fold(0f32, |m, &v| m.max(v.abs()));
9602                eprintln!(
9603                    "layer {li:>3} {:>10} ffn={:<5} rms={rms:>12.4} max={mx:>12.4}",
9604                    match &self.weights.layers[self.phys_layer(li)].attn {
9605                        AttnKind::LinearGdn(_) => "gdn",
9606                        AttnKind::Linear(_) => "vmf",
9607                        AttnKind::ShortConv(_) => "conv",
9608                        _ => "attn",
9609                    },
9610                    match &lw.ffn {
9611                        FfnKind::Moe(_) => "moe",
9612                        FfnKind::Dense(_) => "dense",
9613                        FfnKind::DenseMoe(_) => "dense+moe",
9614                    },
9615                );
9616            }
9617            // Looped Transformer: apply final norm at the end of each loop iteration.
9618            if self.is_loop_end(li) && li + 1 < self.num_layers {
9619                for bi in 0..b {
9620                    let normed = inference::rms_norm(
9621                        &h[bi * hs..(bi + 1) * hs],
9622                        &self.weights.final_norm,
9623                        eps,
9624                        norm_style,
9625                    );
9626                    h[bi * hs..(bi + 1) * hs].copy_from_slice(&normed);
9627                }
9628            }
9629            if std::env::var("CMF_TRACE_H").is_ok() {
9630                let n = h[..hs].iter().map(|v| v.abs()).sum::<f32>() / hs as f32;
9631                let mx = h[..hs].iter().fold(0.0f32, |a, &v| a.max(v.abs()));
9632                eprintln!(
9633                    "layer {li}: mean|h|={n:.4} max|h|={mx:.2} scale={:?}",
9634                    lw.layer_scale
9635                );
9636            }
9637        }
9638        crate::gpu::set_layer(-1); // lm_head/final ops outside layer-split
9639        if prefill_prof_on() {
9640            let ms = |a: &std::sync::atomic::AtomicU64| {
9641                a.load(std::sync::atomic::Ordering::Relaxed) as f64 / 1e6
9642            };
9643            let sp = &attention::ATTN_SPLIT;
9644            let (qkv_enq, qkv_wait, ca_wait, zi) = wgpu_prefill_counters();
9645            eprintln!(
9646                "prefill-split: attention {:.1} ms (proj {:.1} [qkv enqueue {:.1} wait {:.1}, \
9647                 host gate {:.1}], host loop {:.1}, attend {:.1} \
9648                 [q pack {:.1}, device {:.1} of which readback {:.1}], o-proj {:.1}), \
9649                 ffn {:.1} ms (cumulative; zi_mm GEMMs {zi})",
9650                ms(&PREFILL_SPLIT[0]),
9651                ms(&sp[0]),
9652                qkv_enq,
9653                qkv_wait,
9654                ms(&sp[6]),
9655                ms(&sp[1]),
9656                ms(&sp[2]),
9657                ms(&sp[4]),
9658                ms(&sp[5]),
9659                ca_wait,
9660                ms(&sp[3]),
9661                ms(&PREFILL_SPLIT[1]),
9662            );
9663            let [calls, attended, wo, ring, reseed, f16] = wgpu_mirror_counters();
9664            eprintln!(
9665                "prefill-mirror: {calls} calls, {attended} attended ({wo} with O on the card), \
9666                 {ring} ring refusals, {reseed} reseeds, {f16} f16 fallbacks, {} upload \
9667                 fallbacks (cumulative)",
9668                attention::MIRROR_UPLOADS.load(std::sync::atomic::Ordering::Relaxed),
9669            );
9670        }
9671        // A batched span owns a complete set of positions. Publish any
9672        // collecting→sealed transition only after every layer has finished;
9673        // callers that cross into serial/device work must see the new epoch
9674        // before this function returns.
9675        self.o1_progress();
9676        // Every layer of the chunk has appended and attended its rows.
9677        self.swa_trim_tails();
9678        h
9679    }
9680
9681    /// Embed a single token.
9682    fn embed_single(&self, id: u32) -> Vec<f32> {
9683        let mut out = vec![0.0f32; self.hidden_size];
9684        if (id as usize) < self.weights.embed_tokens.rows() {
9685            self.weights.embed_tokens.row_f32(id as usize, &mut out);
9686        }
9687        if self.embed_multiplier != 1.0 {
9688            for v in out.iter_mut() {
9689                *v *= self.embed_multiplier;
9690            }
9691        }
9692        // DeepSeek-V4's hash layers route by TOKEN ID, so the id has to
9693        // reach the forward. It rides in slot 0 (the forward re-reads the
9694        // real embedding itself from the table).
9695        if self.dsv4.is_some()
9696            || self.dsv41.is_some()
9697            || self.qwen4_exp.is_some()
9698        {
9699            let mut v = vec![0.0f32; self.hidden_size.max(1)];
9700            v[0] = id as f32;
9701            return v;
9702        }
9703        // Gemma-3n: the per-layer-embedding half needs the token ID, so
9704        // it rides appended to the embedding; the g3n forward splits it.
9705        if let Some(b) = &self.g3n {
9706            return b.0.extend_embedding(id, &out, self.pool.as_deref());
9707        }
9708        out
9709    }
9710
9711    /// A run of consecutive prefill layers on the GPU for the whole
9712    /// chunk (default-on under CMF_GPU=1; CMF_GPU_CHUNK=0 disables).
9713    /// Eligibility per layer: q8_row weights, plain full attention
9714    /// (no output gate), F32 KV, no o1/masks/gemma extras. Returns the
9715    /// first layer index NOT processed (== `li0` when the run is empty).
9716    #[cfg(target_os = "macos")]
9717    fn chunk_run_gpu(
9718        &mut self,
9719        li0: usize,
9720        h: &mut [f32],
9721        b: usize,
9722        pos0: usize,
9723        embed_ids: Option<&[u32]>,
9724        cap: usize,
9725    ) -> usize {
9726        // (The old streaming attend needed a depth bound at ~1k; the
9727        // GEMM attention scales like the CPU path and lifted it.)
9728        // CMF_GPU_CHUNK=0 disables the graph.
9729        if !crate::gpu::enabled_here()
9730            || std::env::var("CMF_GPU_CHUNK")
9731                .map(|v| v == "0")
9732                .unwrap_or(false)
9733            || b < 32
9734            // Sliding windows with per-layer RoPE ride the chunk graph
9735            // (`metal_graph_swa`: causal_softmax_win, per-layer tables).
9736            || (self.swa.is_some() && !self.metal_graph_swa())
9737            || self.global_attn.is_some()
9738            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
9739            || (self.graph_attn_decline_reason().is_some() && !self.metal_graph_swa())
9740            // Collection owns the exact Q trace and boundary conversion;
9741            // this chunk graph appends dense KV without feeding that trace.
9742            || self.o1_active()
9743            || self.attn_v_norm
9744            || (self.attn_scale - 1.0 / (self.head_dim as f32).sqrt()).abs() > 1e-9
9745        {
9746            return li0;
9747        }
9748        let Some(model) = self.model.clone() else {
9749            return li0;
9750        };
9751        let (nh, nkv, hd, hs) = (
9752            self.num_heads,
9753            self.num_kv_heads,
9754            self.head_dim,
9755            self.hidden_size,
9756        );
9757        // Collect the longest run of consecutive eligible layers.
9758        // Looped Transformer: stop at the loop boundary so the CPU can
9759        // apply loop_final_norm between iterations.
9760        let loop_end = if self.loop_final_norm {
9761            ((li0 / self.physical_layers) + 1) * self.physical_layers
9762        } else {
9763            self.num_layers
9764        };
9765        let mut layers: Vec<crate::gpu_metal::ChunkLayer> = Vec::new();
9766        let mut stored_at: Vec<usize> = Vec::new();
9767        let run_end = self.num_layers.min(loop_end).min(cap);
9768        // Each layer's own RoPE table (Spark-X2.5: sliding and full layers
9769        // differ); the global one for every model without sliding layers.
9770        let tables: Vec<std::sync::Arc<Vec<f32>>> =
9771            (li0..run_end).map(|li| self.layer_inv_freq(li)).collect();
9772        for li in li0..run_end {
9773            let lw = &self.weights.layers[self.phys_layer(li)];
9774            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
9775                break;
9776            }
9777            let AttnKind::Full {
9778                wq,
9779                wk,
9780                wv,
9781                wo,
9782                q_norm,
9783                k_norm,
9784                output_gate: false,
9785                softplus_gate,
9786                bias,
9787            } = &lw.attn
9788            else {
9789                break;
9790            };
9791            // Spark-X2.5's per-head sigmoid g_proj gate, from f32 rows.
9792            let head_gate = match softplus_gate {
9793                None => None,
9794                Some((g, true)) if self.proj_gate_sigmoid => match g.f32_parts() {
9795                    Some((d, r, c)) if r == nh && c == hs => Some(d),
9796                    _ => break,
9797                },
9798                Some(_) => break,
9799            };
9800            let FfnKind::Dense(d) = &lw.ffn else { break };
9801            if !matches!(d.act, Act::Silu | Act::Gelu) || !d.segs.is_empty() {
9802                break;
9803            }
9804            // q8_row (row_scale populated), or q4_tiled / q4tp (row_scale
9805            // empty — their scales are in the payload). Mixing across the
9806            // seven projections of one layer is fine; the encoder branches
9807            // per weight on the tensor's dtype. Anything else refuses.
9808            fn cw(t: &QTensor) -> Option<(usize, usize, usize, &[f32])> {
9809                t.q8_row_parts()
9810                    .or_else(|| t.q4t_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9811                    .or_else(|| t.q4tp_parts().map(|(i, r, c)| (i, r, c, &[][..])))
9812            }
9813            let parts = (
9814                cw(wq),
9815                cw(wk),
9816                cw(wv),
9817                cw(wo),
9818                cw(&d.gate_proj),
9819                cw(&d.up_proj),
9820                cw(&d.down_proj),
9821            );
9822            let (Some(pq), Some(pk), Some(pv), Some(po), Some(pg), Some(pu), Some(pd)) = parts
9823            else {
9824                break;
9825            };
9826            let layer = &self.kv_cache.layers[li];
9827            if layer.mode != crate::kv_cache::KvMode::F32 || layer.o1.is_some() {
9828                break;
9829            }
9830            // A trimmed tail holds every row the chunk's first query reads.
9831            debug_assert!(
9832                layer.base() == 0
9833                    || self
9834                        .layer_window(li)
9835                        .is_some_and(|w| layer.head_len(0) + 1 >= w)
9836            );
9837            stored_at.push(layer.head_len(0));
9838            layers.push(crate::gpu_metal::ChunkLayer {
9839                model: &model,
9840                kv_id: self.graph_kv_id,
9841                layer: li,
9842                wq: pq,
9843                wk: pk,
9844                wv: pv,
9845                wo: po,
9846                gate: pg,
9847                up: pu,
9848                down: pd,
9849                input_norm: &lw.input_norm,
9850                post_norm: &lw.post_norm,
9851                bias: bias
9852                    .as_ref()
9853                    .map(|(a, bb, cc)| (a.as_slice(), bb.as_slice(), cc.as_slice())),
9854                q_norm: q_norm.as_deref(),
9855                k_norm: k_norm.as_deref(),
9856                inv_freq: &tables[li - li0],
9857                rd: self.layer_geom(li).2,
9858                nh,
9859                nkv,
9860                hd,
9861                hs,
9862                inter: d.gate_proj.rows(),
9863                gemma: matches!(self.norm_style, cortiq_core::NormStyle::Gemma),
9864                late_qk_norm: self.qk_norm_after_rope,
9865                eps: self.rms_eps as f32,
9866                window: self.layer_window(li),
9867                head_gate,
9868                gelu: d.act == Act::Gelu,
9869            });
9870        }
9871        if layers.is_empty() {
9872            return li0;
9873        }
9874        let row = nkv * hd;
9875        let mut store: Vec<(Vec<f32>, Vec<f32>, Vec<f32>)> = stored_at
9876            .iter()
9877            .map(|&st| (vec![0f32; b * row], vec![0f32; b * row], vec![0f32; st + b]))
9878            .collect();
9879        let mut io: Vec<crate::gpu_metal::ChunkIo> = Vec::with_capacity(layers.len());
9880        for (i, (ok, ov, oi)) in store.iter_mut().enumerate() {
9881            let li = layers[i].layer;
9882            let layer = &self.kv_cache.layers[li];
9883            io.push(crate::gpu_metal::ChunkIo {
9884                cpu_stored: stored_at[i],
9885                cpu_gen: layer.generation(),
9886                cpu_k: (0..nkv).map(|g| layer.head_keys(g)).collect(),
9887                cpu_v: (0..nkv).map(|g| layer.head_values(g)).collect(),
9888                out_k: ok,
9889                out_v: ov,
9890                imp: oi,
9891            });
9892        }
9893        let n_run = layers.len();
9894        let last = layers.last().map(|l| l.layer + 1).unwrap_or(li0);
9895        // Device-side embedding when the run starts the model and the
9896        // embedding matrix is q8_row-mapped.
9897        let ep = embed_ids.and_then(|ids| {
9898            self.weights
9899                .embed_tokens
9900                .q8_row_parts()
9901                .map(|(idx, rows, _c, rs)| crate::gpu_metal::ChunkEmbed {
9902                    idx,
9903                    rows,
9904                    row_scale: rs,
9905                    ids,
9906                    mult: self.embed_multiplier,
9907                })
9908        });
9909        if embed_ids.is_some() && ep.is_none() {
9910            return li0;
9911        }
9912        if !crate::gpu_metal::chunk_run_gpu(&layers, &mut io, h, b, pos0, ep.as_ref()) {
9913            return li0;
9914        }
9915        drop(io);
9916        drop(layers);
9917        // CPU caches stay the owners of record: append the chunk rows
9918        // and bank the importance masses per layer.
9919        for (i, (ok, ov, oi)) in store.iter().enumerate().take(n_run) {
9920            let li = li0 + i;
9921            let layer = &mut self.kv_cache.layers[li];
9922            for bi in 0..b {
9923                layer.append(
9924                    &ok[bi * row..(bi + 1) * row],
9925                    &ov[bi * row..(bi + 1) * row],
9926                    &[],
9927                );
9928            }
9929            layer.accumulate_imp(oi);
9930        }
9931        last
9932    }
9933
9934    /// Is layer `li` a sliding-window (local-RoPE) layer? Gemma-3:
9935    /// every `pattern`-th layer is global, the rest are local.
9936    fn layer_is_local(&self, li: usize) -> bool {
9937        if let Some(layers) = &self.sliding_layers {
9938            return layers.get(li).copied().unwrap_or(false);
9939        }
9940        match self.swa {
9941            Some((_, pattern)) => (li + 1) % pattern.max(1) != 0,
9942            None => false,
9943        }
9944    }
9945
9946    /// The RoPE table for layer `li` (local layers may have their own;
9947    /// Gemma-4 global layers use the proportional padded table).
9948    fn layer_inv_freq(&self, li: usize) -> std::sync::Arc<Vec<f32>> {
9949        if self.layer_is_local(li) {
9950            if let Some(f) = &self.inv_freq_local {
9951                return f.clone();
9952            }
9953        } else if let Some(f) = &self.inv_freq_global {
9954            return f.clone();
9955        }
9956        self.inv_freq.clone()
9957    }
9958
9959    /// The attend window for layer `li` (None = full context).
9960    fn layer_window(&self, li: usize) -> Option<usize> {
9961        self.swa
9962            .and_then(|(w, _)| self.layer_is_local(li).then_some(w))
9963    }
9964
9965    /// Drop every sliding layer's rows its window can no longer read
9966    /// (`LayerKvCache::trim_window`). Call only at a safe point — after a
9967    /// complete position / pair / chunk walk, never between a layer's
9968    /// appends and its attend. One comparison per layer below the trigger.
9969    ///
9970    /// Off with an MTP head: its verify oracles and rollbacks snapshot and
9971    /// compare per-layer row counts as positions (`CMF_METAL_VERIFY_CHECK`,
9972    /// the MTP caches), which a trimmed tail would break.
9973    fn swa_trim_tails(&mut self) {
9974        let Some((slack, align)) = self.swa_trim else {
9975            return;
9976        };
9977        if self.mtp.is_some() || self.mimo_mtp.is_some() || !self.dsv4_mtp.is_empty() {
9978            return;
9979        }
9980        for li in 0..self.num_layers.min(self.kv_cache.layers.len()) {
9981            if let Some(w) = self.layer_window(li) {
9982                self.kv_cache.layers[li].trim_window(w, slack, align);
9983            }
9984        }
9985    }
9986
9987    fn layer_num_heads(&self, li: usize) -> usize {
9988        self.attention_heads_per_layer
9989            .as_ref()
9990            .and_then(|v| v.get(li).copied())
9991            .unwrap_or(self.num_heads)
9992    }
9993
9994    fn layer_rope_scale(&self, li: usize) -> f32 {
9995        if self.layer_is_local(li) {
9996            self.rope_scale_local
9997        } else {
9998            self.rope_scale
9999        }
10000    }
10001
10002    /// Attention geometry of layer `li`: (num_kv_heads, head_dim,
10003    /// rotary_dim). Gemma-4 global layers override all three.
10004    fn layer_geom(&self, li: usize) -> (usize, usize, usize) {
10005        if !self.layer_is_local(li) {
10006            if let Some((ghd, gkv)) = self.global_attn {
10007                return (gkv, ghd, ghd);
10008            }
10009        }
10010        (
10011            self.layer_num_kv_heads(li),
10012            self.head_dim,
10013            if self.layer_is_local(li) {
10014                self.rotary_dim_local.unwrap_or(self.rotary_dim)
10015            } else {
10016                self.rotary_dim
10017            },
10018        )
10019    }
10020
10021    /// KV heads of layer `li` (virtual index): the per-layer count when the
10022    /// model has one (MiMo-V2), else the uniform `num_kv_heads`.
10023    fn layer_num_kv_heads(&self, li: usize) -> usize {
10024        self.kv_heads_per_layer
10025            .as_ref()
10026            .and_then(|v| v.get(self.phys_layer(li)).copied())
10027            .unwrap_or(self.num_kv_heads)
10028    }
10029
10030    /// V head width of layer `li` (≤ its head_dim).
10031    fn layer_v_dim(&self, li: usize) -> usize {
10032        let (_, hd, _) = self.layer_geom(li);
10033        self.v_head_dim.unwrap_or(hd).min(hd)
10034    }
10035
10036    /// Install a per-layer KV geometry: KV heads per PHYSICAL layer and/or
10037    /// a V head width narrower than `head_dim` (MiMo-V2). Validates it and
10038    /// reshapes the caches of every layer whose KV head count differs from
10039    /// `num_kv_heads`. The loader and the tests share this one path, so a
10040    /// hand-built pipeline cannot hold a geometry the loader would refuse.
10041    /// Call before the first forward (it drops cached rows of reshaped
10042    /// layers). Refuses combinations whose paths would read it wrong:
10043    /// Gemma-4 global layers and MLA carry their own geometry.
10044    pub fn set_attn_geometry(
10045        &mut self,
10046        kv_heads_per_layer: Option<Vec<usize>>,
10047        v_head_dim: Option<usize>,
10048    ) -> Result<(), String> {
10049        if kv_heads_per_layer.is_some() || v_head_dim.is_some() {
10050            if self.global_attn.is_some() {
10051                return Err(
10052                    "per-layer KV heads / v_head_dim cannot combine with Gemma-4 global \
10053                     attention geometry"
10054                        .into(),
10055                );
10056            }
10057            if self
10058                .weights
10059                .layers
10060                .iter()
10061                .any(|lw| matches!(lw.attn, AttnKind::Mla(_)))
10062            {
10063                return Err("per-layer KV heads / v_head_dim cannot combine with MLA".into());
10064            }
10065        }
10066        if let Some(vd) = v_head_dim {
10067            if vd == 0 || vd > self.head_dim {
10068                return Err(format!(
10069                    "v_head_dim {vd} must be in 1..={} (head_dim)",
10070                    self.head_dim
10071                ));
10072            }
10073        }
10074        if let Some(v) = &kv_heads_per_layer {
10075            if v.len() != self.physical_layers {
10076                return Err(format!(
10077                    "kv_heads_per_layer has {} entries, expected {} layers",
10078                    v.len(),
10079                    self.physical_layers
10080                ));
10081            }
10082            for (li, &nkv) in v.iter().enumerate() {
10083                let is_attn = matches!(
10084                    self.weights.layers.get(li).map(|lw| &lw.attn),
10085                    Some(AttnKind::Full { .. }) | None
10086                );
10087                if !is_attn {
10088                    continue;
10089                }
10090                let nh = self
10091                    .attention_heads_per_layer
10092                    .as_ref()
10093                    .and_then(|h| h.get(li).copied())
10094                    .unwrap_or(self.num_heads);
10095                if nkv == 0 || nh % nkv != 0 {
10096                    return Err(format!(
10097                        "layer {li}: {nkv} KV heads must be nonzero and divide {nh} Q heads"
10098                    ));
10099                }
10100            }
10101        }
10102        self.kv_heads_per_layer = kv_heads_per_layer;
10103        self.v_head_dim = v_head_dim.filter(|&vd| vd != self.head_dim);
10104        if self.kv_heads_per_layer.is_some() {
10105            for li in 0..self.kv_cache.layers.len() {
10106                let full = matches!(
10107                    self.weights
10108                        .layers
10109                        .get(self.phys_layer(li))
10110                        .map(|lw| &lw.attn),
10111                    Some(AttnKind::Full { .. })
10112                );
10113                let nkv = self.layer_num_kv_heads(li);
10114                let cache = &self.kv_cache.layers[li];
10115                if full && (cache.num_kv_heads != nkv || cache.head_dim != self.head_dim) {
10116                    let sinks = cache.sinks.clone();
10117                    self.kv_cache.layers[li] =
10118                        crate::kv_cache::LayerKvCache::new(nkv, self.head_dim);
10119                    self.kv_cache.layers[li].sinks = sinks;
10120                }
10121            }
10122        }
10123        Ok(())
10124    }
10125
10126    /// Attach learned attention-sink logits (one per Q head) to PHYSICAL
10127    /// layer `phys` — every virtual layer that runs it. The loader calls
10128    /// this for each `model.layers.N.self_attn.sinks` tensor.
10129    pub fn set_layer_sinks(&mut self, phys: usize, sinks: Vec<f32>) -> Result<(), String> {
10130        let Some(lw) = self.weights.layers.get(phys) else {
10131            return Err(format!("sinks for layer {phys}: no such layer"));
10132        };
10133        if !matches!(lw.attn, AttnKind::Full { .. }) {
10134            return Err(format!(
10135                "sinks for layer {phys}: only softmax (Full) attention layers take sinks"
10136            ));
10137        }
10138        let nh = self
10139            .attention_heads_per_layer
10140            .as_ref()
10141            .and_then(|h| h.get(phys).copied())
10142            .unwrap_or(self.num_heads);
10143        if sinks.len() != nh {
10144            return Err(format!(
10145                "sinks for layer {phys}: {} values, expected one per Q head ({nh})",
10146                sinks.len()
10147            ));
10148        }
10149        if let Some(bad) = sinks.iter().find(|v| !v.is_finite()) {
10150            return Err(format!("sinks for layer {phys}: non-finite value {bad}"));
10151        }
10152        for li in 0..self.kv_cache.layers.len() {
10153            if self.phys_layer(li) == phys {
10154                self.kv_cache.layers[li].sinks = Some(sinks.clone());
10155            }
10156        }
10157        Ok(())
10158    }
10159
10160    /// Why the GPU attention graphs cannot serve this model, if they
10161    /// cannot: the wgpu whole-token and batched graphs, the greedy
10162    /// multi-burst, the q1 attention dropin and the Metal block/chunk/rows
10163    /// graphs all assume ONE (num_kv_heads, head_dim) geometry, V heads as
10164    /// wide as K, a single RoPE table, full-context attention and a plain
10165    /// softmax. A model outside that contract runs on the CPU layer walk
10166    /// (and the per-op GPU matvecs) until a graph learns it — never on a
10167    /// graph that would read it wrong. None = no attention-level reason
10168    /// (the graph builders still check weights and layer kinds).
10169    pub fn graph_attn_decline_reason(&self) -> Option<&'static str> {
10170        if self.kv_heads_per_layer.is_some() {
10171            return Some("per-layer KV head counts");
10172        }
10173        if self.v_head_dim.is_some_and(|vd| vd != self.head_dim) {
10174            return Some("V heads narrower than Q/K heads");
10175        }
10176        if self.kv_cache.layers.iter().any(|l| l.sinks.is_some()) {
10177            return Some("learned attention sinks");
10178        }
10179        if self.swa.is_some() || self.sliding_layers.is_some() {
10180            return Some("sliding-window layers");
10181        }
10182        None
10183    }
10184
10185    /// Can the Metal block graph carry this model's sliding-window layers?
10186    /// Both of its attention forms take a layer's own window, RoPE table
10187    /// and rotary width: the device attend (`AttnDeviceParams::window`, the
10188    /// layer's `inv_freq`/`rd`) and the sandwich's host attend (exactly the
10189    /// CPU path's attention). True when the sliding window is the model's
10190    /// only attention-level decline (Spark-X2.5); per-layer KV heads,
10191    /// narrow V, sinks, Gemma-4 global geometry and scaled RoPE positions
10192    /// still decline.
10193    ///
10194    /// Per-layer Q heads (Laguna) and capped scores (Gemma-2) decline here
10195    /// too, with V norm (Gemma-4): the chunk prefill's own gate checks
10196    /// neither of the first two, and before this door opened its `swa`
10197    /// refusal kept every sliding model with them off the chunk graph.
10198    #[cfg(target_os = "macos")]
10199    fn metal_graph_swa(&self) -> bool {
10200        (self.swa.is_some() || self.sliding_layers.is_some())
10201            && self.attention_heads_per_layer.is_none()
10202            && !self.attn_v_norm
10203            && self.attn_softcap == 0.0
10204            && self.kv_heads_per_layer.is_none()
10205            && self.v_head_dim.map_or(true, |vd| vd == self.head_dim)
10206            && !self.kv_cache.layers.iter().any(|l| l.sinks.is_some())
10207            && self.global_attn.is_none()
10208            && self.inv_freq_global.is_none()
10209            && (0..self.num_layers).all(|li| self.layer_rope_scale(li) == 1.0)
10210            && !(0..self.num_layers).any(|li| {
10211                self.layer_is_local(li)
10212                    && self.inv_freq_local.is_none()
10213                    && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
10214            })
10215    }
10216
10217    /// Why the WGPU graphs (whole-token, batched prefill, greedy burst)
10218    /// cannot run this model's attention, if they cannot. Per-layer KV
10219    /// heads, V narrower than K, learned sinks and sliding windows ride
10220    /// their per-layer geometry (`GraphAttnGeom`, the ATTEND_X kernels);
10221    /// what that geometry does not express keeps the decline, by name.
10222    /// None for every model with one attention geometry.
10223    pub fn wgpu_graph_attn_decline(&self) -> Option<&'static str> {
10224        self.graph_attn_decline_reason()?;
10225        if self.global_attn.is_some() {
10226            return Some("per-layer head width (Gemma-4 global layers) with per-layer geometry");
10227        }
10228        if self.attention_heads_per_layer.is_some() {
10229            return Some("per-layer Q head counts with per-layer geometry");
10230        }
10231        if self.attn_v_norm {
10232            return Some("V norm with per-layer geometry");
10233        }
10234        if (0..self.num_layers).any(|li| self.layer_rope_scale(li) != 1.0) {
10235            return Some("scaled RoPE positions with per-layer geometry");
10236        }
10237        if self.weights.layers.iter().any(|lw| {
10238            lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some()
10239        }) {
10240            return Some("sandwich norms / layer scale with per-layer geometry");
10241        }
10242        if self.weights.layers.iter().any(|lw| {
10243            matches!(
10244                &lw.attn,
10245                AttnKind::Full {
10246                    output_gate: true,
10247                    ..
10248                }
10249            )
10250        }) && self.v_head_dim.is_some()
10251        {
10252            return Some("gated attention with V narrower than K");
10253        }
10254        if (0..self.num_layers).any(|li| {
10255            self.layer_is_local(li)
10256                && self.inv_freq_local.is_none()
10257                && self.rotary_dim_local.is_some_and(|r| r != self.rotary_dim)
10258        }) {
10259            return Some("local rotary width without a local RoPE table");
10260        }
10261        None
10262    }
10263
10264    /// The wgpu graphs' attention geometry for layer `li` (virtual index):
10265    /// Some only for a model whose layers do not share one (MiMo-V2) — KV
10266    /// heads, V width, rotary width and RoPE table, window and sinks of
10267    /// THIS layer, exactly what the CPU attention reads for it.
10268    fn graph_attn_geom(&self, li: usize) -> Option<crate::gpu::GraphAttnGeom<'_>> {
10269        self.graph_attn_decline_reason()?;
10270        let (nkv, _hd, rd) = self.layer_geom(li);
10271        let invf: &[f32] = if self.layer_is_local(li) {
10272            match &self.inv_freq_local {
10273                Some(f) => f.as_slice(),
10274                None => self.inv_freq.as_slice(),
10275            }
10276        } else {
10277            match &self.inv_freq_global {
10278                Some(f) => f.as_slice(),
10279                None => self.inv_freq.as_slice(),
10280            }
10281        };
10282        Some(crate::gpu::GraphAttnGeom {
10283            nkv,
10284            dv: self.layer_v_dim(li),
10285            rd,
10286            invf,
10287            window: self.layer_window(li),
10288            sink: self.kv_cache.layers[li].sinks.as_deref(),
10289        })
10290    }
10291
10292    /// The decode graph's device K/V mirror that layer `li`'s batched
10293    /// prefill appends its chunk to and attends against (wgpu), instead
10294    /// of uploading the layer's whole K/V prefix every chunk and the whole
10295    /// cache again at the first decode token. Only where the token graph
10296    /// will read that very mirror: the decode graph on, not refused, and
10297    /// this layer's geometry what it requests (`graph_attn_geom`: KV heads,
10298    /// V as wide as K, the window as a ring, no sink).
10299    ///
10300    /// Spark-X2.5 only (`proj_gate_sigmoid`), the one architecture whose
10301    /// output was measured identical on it. `CMF_PREFILL_MIRROR=0` keeps
10302    /// the host upload (A/B).
10303    #[cfg(feature = "gpu")]
10304    fn prefill_mirror_target(&self, li: usize) -> Option<crate::gpu::PrefillMirror> {
10305        if !self.proj_gate_sigmoid
10306            || std::env::var("CMF_PREFILL_MIRROR").as_deref() == Ok("0")
10307            || !crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
10308            || self.graph_refused()
10309            || self.o1_active()
10310            || self.attn_softcap > 0.0
10311            || self.wgpu_graph_attn_decline().is_some()
10312        {
10313            return None;
10314        }
10315        let g = self.graph_attn_geom(li)?;
10316        let (nkv, hd, _) = self.layer_geom(li);
10317        // The token graph passes the call-wide head width (layer 0's).
10318        if g.nkv != nkv || g.dv != hd || g.sink.is_some() || hd != self.layer_geom(0).1 {
10319            return None;
10320        }
10321        Some(crate::gpu::PrefillMirror {
10322            kv_id: self.graph_kv_id,
10323            layer: li,
10324            limit: self.kv_cache.max_seq_len,
10325        })
10326    }
10327
10328    /// Bring the host KV cache of every Full-attention layer in
10329    /// `[from, upto)` up to `position` rows from the wgpu mirrors, where a
10330    /// device graph advanced a layer that the host is about to run: a
10331    /// device prefix that shrank since the prompt (or a batched prefill
10332    /// prefix longer than the decode one). A layer whose mirror does not
10333    /// hold the missing rows is left alone. Rows a sliding layer's ring
10334    /// no longer holds come back as zeros — outside every window that
10335    /// will read them.
10336    #[cfg(feature = "gpu")]
10337    fn pull_lagging_host_kv(&mut self, from: usize, upto: usize, position: usize) {
10338        let kv_id = self.graph_kv_id;
10339        for li in from..upto.min(self.num_layers) {
10340            if !matches!(
10341                self.weights.layers[self.phys_layer(li)].attn,
10342                AttnKind::Full { .. }
10343            ) {
10344                continue;
10345            }
10346            // Absolute: a trimmed sliding tail holds positions base.. only.
10347            let host = self.kv_cache.layers[li].pos_len();
10348            if host >= position {
10349                continue;
10350            }
10351            let Some(dev) = crate::gpu::graph_kv_stored(kv_id, li) else {
10352                continue;
10353            };
10354            let to = dev.min(position);
10355            if to <= host {
10356                continue;
10357            }
10358            let (nkv, hd) = {
10359                let c = &self.kv_cache.layers[li];
10360                (c.num_kv_heads, c.head_dim)
10361            };
10362            let Some((k, v, first_valid)) =
10363                crate::gpu::graph_kv_pull_host(kv_id, li, host, to, nkv, hd)
10364            else {
10365                continue;
10366            };
10367            // A sliding layer only ever reads its last `window` rows; a
10368            // full-context layer needs every row it did not have.
10369            let need_from = match self.layer_window(li) {
10370                Some(w) => host.max((position + 1).saturating_sub(w)),
10371                None => host,
10372            };
10373            if first_valid > need_from {
10374                tracing::warn!(
10375                    "layer {li}: device KV rows {host}..{to} no longer resident \
10376                     (from {first_valid}); host attention will miss them"
10377                );
10378            }
10379            let row = nkv * hd;
10380            let cache = &mut self.kv_cache.layers[li];
10381            for p in 0..to - host {
10382                cache.append(&k[p * row..(p + 1) * row], &v[p * row..(p + 1) * row], &[]);
10383            }
10384        }
10385    }
10386
10387    /// Log (once per graph site and pipeline) that `site` declined for
10388    /// `reason`. The lines are kept so a caller or a test can read them.
10389    fn note_graph_decline(&self, site: &'static str, reason: &'static str) {
10390        let mut seen = self.graph_declines.borrow_mut();
10391        if !seen.iter().any(|&(s, r)| s == site && r == reason) {
10392            tracing::warn!("{site} declined: {reason} (CPU attention path)");
10393            seen.push((site, reason));
10394        }
10395    }
10396
10397    /// The GPU-graph declines this pipeline has logged so far, as the
10398    /// logged lines.
10399    pub fn graph_declines(&self) -> Vec<String> {
10400        self.graph_declines
10401            .borrow()
10402            .iter()
10403            .map(|(site, reason)| format!("{site} declined: {reason} (CPU attention path)"))
10404            .collect()
10405    }
10406
10407    /// Does layer `li` have the plain attention geometry the historical
10408    /// head-masked f32 path (`multi_head_attention`) assumes — pipeline-wide
10409    /// KV heads / head_dim / RoPE table, full context, no sink, V as wide
10410    /// as K? Anything else runs the dense `qwen_attention` instead.
10411    /// `CMF_LAYER_DUMP` writer (see `Pipeline::layer_dump`): one position's
10412    /// hidden after layer `li` as raw little-endian f32 into
10413    /// `<dir>/p{pos:06}_l{li:02}.f32`. A failed write is reported once and
10414    /// never stops the forward — the dump is a diagnostic.
10415    fn dump_layer_row(&self, pos: usize, li: usize, row: &[f32]) {
10416        let Some(dir) = &self.layer_dump else {
10417            return;
10418        };
10419        let mut bytes = Vec::with_capacity(row.len() * 4);
10420        for v in row {
10421            bytes.extend_from_slice(&v.to_le_bytes());
10422        }
10423        let path = dir.join(format!("p{pos:06}_l{li:02}.f32"));
10424        if let Err(e) = std::fs::create_dir_all(dir).and_then(|_| std::fs::write(&path, &bytes)) {
10425            use std::sync::atomic::{AtomicBool, Ordering};
10426            static SAID: AtomicBool = AtomicBool::new(false);
10427            if !SAID.swap(true, Ordering::Relaxed) {
10428                tracing::error!("CMF_LAYER_DUMP: cannot write {}: {e}", path.display());
10429            }
10430        }
10431    }
10432
10433    /// Decide the MiMo-V2 expert placement once (`crate::mimo_moe`). Any
10434    /// other model turns the slot off on the first call.
10435    fn mimo_moe_prepare(&mut self) {
10436        if !self.mimo_moe.is_undecided() {
10437            return;
10438        }
10439        let slot = {
10440            let layers: Vec<(usize, &MoeFfn)> = (0..self.num_layers)
10441                .filter_map(
10442                    |li| match &self.weights.layers.get(self.phys_layer(li))?.ffn {
10443                        FfnKind::Moe(m) => Some((li, m)),
10444                        _ => None,
10445                    },
10446                )
10447                .collect();
10448            // One bank lives on one device: an in-process multi-GPU split
10449            // keeps the whole-layer path.
10450            if layers.is_empty()
10451                || self.physical_layers != self.num_layers
10452                || self.gpu_plan.is_some()
10453            {
10454                crate::mimo_moe::Slot::Off
10455            } else {
10456                // Whether a whole-token graph could run this model's layers
10457                // (then a whole-layer prefix is one submit, not per-layer
10458                // fences).
10459                let graph_prefix =
10460                    self.wgpu_graph_attn_decline().is_none() && crate::gpu::wgpu_graph_default();
10461                crate::mimo_moe::Slot::decide(&layers, self.num_layers, graph_prefix)
10462            }
10463        };
10464        self.mimo_moe = slot;
10465    }
10466
10467    #[cfg(test)]
10468    pub(crate) fn test_graph_kv_id(&self) -> u64 {
10469        self.graph_kv_id
10470    }
10471
10472    /// Dynamic MiMo layer: one device attention graph, followed by a bank
10473    /// frame. Both decode and short verification use this same attention
10474    /// path and absolute layer key; the host KV may intentionally lag.
10475    pub(crate) fn mimo_graph_layer_rows(
10476        &mut self,
10477        li: usize,
10478        h: &mut [f32],
10479        positions: &[usize],
10480    ) -> crate::gpu::BatchGraphOutcome {
10481        use crate::gpu::BatchGraphOutcome as Out;
10482        let b = positions.len();
10483        if !(1..=4).contains(&b)
10484            || h.len() != b * self.hidden_size
10485            || !self.mimo_moe.is_dynamic(li, true)
10486            || !crate::gpu::enabled_here()
10487            || !crate::gpu::wgpu_active()
10488            || self.o1_active()
10489            || self.physical_layers != self.num_layers
10490            // The pair-fusion diagnostic (and an explicit graph-off run)
10491            // rewinds only host KV. A hidden singleton attention graph here
10492            // would leave device mirrors ahead of the next host position.
10493            || std::env::var("CMF_GPU_WGPU_GRAPH").as_deref() == Ok("0")
10494            || std::env::var("CMF_MIMO_ATTN_GRAPH").as_deref() == Ok("0")
10495            || self.wgpu_graph_attn_decline().is_some()
10496        {
10497            return Out::Declined;
10498        }
10499        let attn_started = std::time::Instant::now();
10500        let outcome = {
10501            let lw = &self.weights.layers[li];
10502            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
10503                return Out::Declined;
10504            }
10505            let FfnKind::Moe(m) = &lw.ffn else {
10506                return Out::Declined;
10507            };
10508            let AttnKind::Full {
10509                wq,
10510                wk,
10511                wv,
10512                wo,
10513                q_norm,
10514                k_norm,
10515                output_gate,
10516                softplus_gate,
10517                bias,
10518            } = &lw.attn
10519            else {
10520                return Out::Declined;
10521            };
10522            if *output_gate || softplus_gate.is_some() {
10523                return Out::Declined;
10524            }
10525            let Some(model) = wq.graph_weight().map(|(m, ..)| m.clone()).or_else(|| {
10526                m.experts
10527                    .first()?
10528                    .gate_proj
10529                    .mapped_q4tp()
10530                    .map(|(m, _)| m.clone())
10531            }) else {
10532                return Out::Declined;
10533            };
10534            fn gw<'a>(
10535                t: &'a QTensor,
10536                owner: &std::sync::Arc<cortiq_core::CmfModel>,
10537            ) -> Option<crate::gpu::GraphW<'a>> {
10538                if let Some((m, idx, kind, rs)) = t.graph_weight() {
10539                    if m.uid() != owner.uid() || t.has_prism_contract() {
10540                        return None;
10541                    }
10542                    return Some(crate::gpu::GraphW {
10543                        idx,
10544                        kind,
10545                        row_scale: rs,
10546                        data: &[],
10547                        prism: crate::gpu::GraphPrismOp::None,
10548                        affine: false,
10549                    });
10550                }
10551                t.as_f32().map(|data| crate::gpu::GraphW {
10552                    idx: 0,
10553                    kind: 4,
10554                    row_scale: &[],
10555                    data,
10556                    prism: crate::gpu::GraphPrismOp::None,
10557                    affine: false,
10558                })
10559            }
10560            let (Some(q), Some(k), Some(v), Some(o)) = (
10561                gw(wq, &model),
10562                gw(wk, &model),
10563                gw(wv, &model),
10564                gw(wo, &model),
10565            ) else {
10566                return Out::Declined;
10567            };
10568            let layer = crate::gpu::GraphLayer {
10569                input_norm: &lw.input_norm,
10570                post_norm: &lw.post_norm,
10571                ffn: crate::gpu::GraphFfn::AttentionOnly,
10572                attn: crate::gpu::GraphAttn::Full {
10573                    wq: q,
10574                    wk: k,
10575                    wv: v,
10576                    wo: o,
10577                    q_norm: q_norm.as_deref(),
10578                    k_norm: k_norm.as_deref(),
10579                    late_qk_norm: self.qk_norm_after_rope,
10580                    bias: bias
10581                        .as_ref()
10582                        .map(|(q, k, v)| (q.as_slice(), k.as_slice(), v.as_slice())),
10583                    output_gate: false,
10584                    cpu_k: self.kv_cache.layers[li].k_heads(),
10585                    cpu_v: self.kv_cache.layers[li].v_heads(),
10586                    cpu_base: self.kv_cache.layers[li].base(),
10587                    geom: self.graph_attn_geom(li),
10588                    head_gate: None,
10589                },
10590            };
10591            let (nkv, hd, rd) = self.layer_geom(li);
10592            crate::gpu::forward_batch_graph_at(
10593                &model,
10594                self.graph_kv_id,
10595                li,
10596                &[layer],
10597                &self.inv_freq,
10598                h,
10599                self.layer_num_heads(li),
10600                nkv,
10601                hd,
10602                rd,
10603                self.hidden_size,
10604                1,
10605                positions,
10606                self.kv_cache.max_seq_len,
10607                self.norm_style == cortiq_core::NormStyle::Gemma,
10608                self.rms_eps as f32,
10609                self.attn_scale,
10610                b,
10611                &[],
10612                self.o1_epoch,
10613                None,
10614                None,
10615            )
10616        };
10617        match outcome {
10618            Out::Completed => {}
10619            Out::Declined => return Out::Declined,
10620            Out::Failed => {
10621                self.graph_failed
10622                    .store(true, std::sync::atomic::Ordering::Relaxed);
10623                return Out::Failed;
10624            }
10625        }
10626        let attn_ns = attn_started.elapsed().as_nanos() as u64;
10627        let hs = self.hidden_size;
10628        let lw = &self.weights.layers[li];
10629        let FfnKind::Moe(m) = &lw.ffn else {
10630            unreachable!()
10631        };
10632        let mut post = vec![0.0; h.len()];
10633        for (x, y) in h.chunks_exact(hs).zip(post.chunks_exact_mut(hs)) {
10634            inference::rms_norm_into(x, &lw.post_norm, self.rms_eps, self.norm_style, y);
10635        }
10636        let mut ffn = if b == 1 {
10637            moe_ffn_banked(&mut self.mimo_moe, li, m, &post, self.pool.as_deref())
10638        } else {
10639            moe_ffn_banked_rows(
10640                &mut self.mimo_moe,
10641                li,
10642                m,
10643                &post,
10644                b,
10645                hs,
10646                self.pool.as_deref(),
10647            )
10648        };
10649        for (x, &f) in h.iter_mut().zip(&ffn) {
10650            *x += f;
10651        }
10652        attention::recycle_buf(&mut ffn);
10653        if self.layer_dump.is_some() {
10654            for (&pos, row) in positions.iter().zip(h.chunks_exact(hs)) {
10655                self.dump_layer_row(pos, li, row);
10656            }
10657        }
10658        crate::mimo_moe::note_attention_graph(b, attn_ns);
10659        Out::Completed
10660    }
10661
10662    fn layer_attn_plain(&self, li: usize) -> bool {
10663        self.kv_heads_per_layer.is_none()
10664            && self.v_head_dim.is_none()
10665            && self.global_attn.is_none()
10666            && self.layer_window(li).is_none()
10667            && self.kv_cache.layers[li].sinks.is_none()
10668    }
10669
10670    /// Forward one position through all layers (hybrid dispatch).
10671    fn forward_layers(
10672        &mut self,
10673        hidden: &[f32],
10674        position: usize,
10675        task_mask: Option<&TaskMask>,
10676    ) -> Vec<f32> {
10677        let out = self.forward_layers_upto(hidden, position, task_mask, None);
10678        self.o1_progress();
10679        self.swa_trim_tails();
10680        out
10681    }
10682
10683    // ── Network pipeline-split building blocks (coordinator/worker) ──
10684    // A remote worker owns layers [from ..= upto] and their KV; the
10685    // coordinator owns the rest plus embed / final norm / head. Attention
10686    // causality is per-layer, so a whole prompt's boundary hiddens ship
10687    // as one batch and decode ships one vector per token.
10688
10689    /// Embed one token id (embed multiplier applied).
10690    pub fn embed_id(&self, id: u32) -> Vec<f32> {
10691        self.embed_single(id)
10692    }
10693
10694    /// Refuse the archs/modes whose forward cannot be cut at a layer
10695    /// boundary. Loud by design: a split that silently changed the math
10696    /// would be a chimera.
10697    pub fn split_supported(&self) -> Result<(), String> {
10698        if self.dsv4.is_some() {
10699            return Err(
10700                "network split: DeepSeek-V4 runs its own fused stack (not splittable yet)".into(),
10701            );
10702        }
10703        if self.dsv41.is_some() {
10704            return Err(
10705                "network split: DeepSeek-V4.1 owns the shared CED/CSA2 state (not splittable)"
10706                    .into(),
10707            );
10708        }
10709        if self.qwen4_exp.is_some() {
10710            return Err(
10711                "network split: Qwen3.8-Flash-Next hyper/QSA stack is not splittable yet".into(),
10712            );
10713        }
10714        if self.g3n.is_some() {
10715            return Err(
10716                "network split: Gemma-3n runs its own AltUp stack (not splittable yet)".into(),
10717            );
10718        }
10719        Ok(())
10720    }
10721
10722    /// Forward `hidden` through layers [from ..= upto] at `position`,
10723    /// appending those layers' KV/state. Both split sides call this
10724    /// over their own range; a task mask applies to the span's own
10725    /// layers (each side masks what it runs).
10726    pub fn forward_span(
10727        &mut self,
10728        hidden: &[f32],
10729        position: usize,
10730        from: usize,
10731        upto: usize,
10732        task_mask: Option<&TaskMask>,
10733    ) -> Result<Vec<f32>, String> {
10734        self.split_supported()?;
10735        if from > upto || upto >= self.num_layers {
10736            return Err(format!(
10737                "forward_span: layer range {from}..={upto} outside 0..{}",
10738                self.num_layers
10739            ));
10740        }
10741        if hidden.len() != self.hidden_size {
10742            return Err(format!(
10743                "forward_span: hidden len {} ≠ hidden_size {}",
10744                hidden.len(),
10745                self.hidden_size
10746            ));
10747        }
10748        let out = self.forward_layers_span(hidden, position, task_mask, from, Some(upto));
10749        self.o1_progress();
10750        self.swa_trim_tails();
10751        if self
10752            .graph_failed
10753            .swap(false, std::sync::atomic::Ordering::Relaxed)
10754        {
10755            self.cancel
10756                .store(false, std::sync::atomic::Ordering::Relaxed);
10757            self.clear_sequence_state();
10758            return Err("forward_span: deferred O(1) transition failed".into());
10759        }
10760        Ok(out)
10761    }
10762
10763    /// Final norm + lm_head over a boundary hidden (the final-logit
10764    /// softcap is applied by lm_head_forward itself).
10765    pub fn logits_from_hidden(&mut self, hidden: &[f32]) -> Vec<f32> {
10766        let normed = inference::rms_norm(
10767            hidden,
10768            &self.weights.final_norm,
10769            self.rms_eps,
10770            self.norm_style,
10771        );
10772        self.lm_head_forward(&normed)
10773    }
10774
10775    /// Sample the next token with this pipeline's sampler state.
10776    pub fn sample_next(&mut self, logits: &[f32], past_tokens: &[u32]) -> u32 {
10777        sampler::sample_with_scratch(
10778            logits,
10779            &self.sampler_config,
10780            past_tokens,
10781            &mut self.rng,
10782            &mut self.sampler_scratch,
10783        )
10784    }
10785
10786    /// Fresh sequence: clear KV, reuse history and device mirrors.
10787    pub fn reset_session(&mut self) {
10788        self.clear_sequence_state();
10789    }
10790
10791    /// Batched span prefill from token ids (coordinator side): embed +
10792    /// layers [0 ..= upto]; returns the boundary hiddens of ALL positions
10793    /// (ids.len() × hidden). Rides the same layer-major machinery as the
10794    /// local prefill; falls back to the per-position walk under
10795    /// CMF_PREFILL=seq.
10796    pub fn prefill_span_ids(
10797        &mut self,
10798        ids: &[u32],
10799        start_pos: usize,
10800        upto: usize,
10801        task_mask: Option<&TaskMask>,
10802    ) -> Result<Vec<f32>, String> {
10803        self.split_supported()?;
10804        if upto >= self.num_layers {
10805            return Err(format!(
10806                "prefill_span_ids: upto {upto} outside 0..{}",
10807                self.num_layers
10808            ));
10809        }
10810        // Same predicate as the whole-stack prefill: a span whose GDN
10811        // state lives on the device must walk positions through the
10812        // graph, not through the batched CPU span.
10813        if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10814            let out =
10815                self.prefill_batch_span(PrefillIn::Ids(ids), start_pos, task_mask, 0, upto + 1);
10816            self.check_o1_progress_failure("prefill_span_ids")?;
10817            Ok(out)
10818        } else {
10819            let hs = self.hidden_size;
10820            let mut out = Vec::with_capacity(ids.len() * hs);
10821            for (i, &id) in ids.iter().enumerate() {
10822                let emb = self.embed_id(id);
10823                out.extend_from_slice(&self.forward_span(
10824                    &emb,
10825                    start_pos + i,
10826                    0,
10827                    upto,
10828                    task_mask,
10829                )?);
10830            }
10831            Ok(out)
10832        }
10833    }
10834
10835    /// Batched span prefill from boundary hiddens (worker side): layers
10836    /// [from ..= upto] for every position in the batch; returns the batch.
10837    pub fn prefill_span_hidden(
10838        &mut self,
10839        hidden: &[f32],
10840        start_pos: usize,
10841        from: usize,
10842        upto: usize,
10843        task_mask: Option<&TaskMask>,
10844    ) -> Result<Vec<f32>, String> {
10845        self.split_supported()?;
10846        let hs = self.hidden_size;
10847        if hidden.is_empty() || hidden.len() % hs != 0 {
10848            return Err(format!(
10849                "prefill_span_hidden: {} floats is not a multiple of hidden {hs}",
10850                hidden.len()
10851            ));
10852        }
10853        if from > upto || upto >= self.num_layers {
10854            return Err(format!(
10855                "prefill_span_hidden: layer range {from}..={upto} outside 0..{}",
10856                self.num_layers
10857            ));
10858        }
10859        if self.can_prefill_batched() && !self.graph_prefill_preferred() {
10860            let out = self.prefill_batch_span(
10861                PrefillIn::Hidden(hidden),
10862                start_pos,
10863                task_mask,
10864                from,
10865                upto + 1,
10866            );
10867            self.check_o1_progress_failure("prefill_span_hidden")?;
10868            Ok(out)
10869        } else {
10870            let b = hidden.len() / hs;
10871            let mut out = Vec::with_capacity(hidden.len());
10872            for i in 0..b {
10873                let h = self.forward_span(
10874                    &hidden[i * hs..(i + 1) * hs],
10875                    start_pos + i,
10876                    from,
10877                    upto,
10878                    task_mask,
10879                )?;
10880                out.extend_from_slice(&h);
10881            }
10882            Ok(out)
10883        }
10884    }
10885
10886    /// Build the whole-token wgpu graph for a pure-attention q1 model (every
10887    /// layer Full q1 + dense q1 FFN, no gate/bias). Returns the post-stack
10888    /// hidden (caller does final norm + lm_head), or None to fall back.
10889    fn try_token_graph_wgpu(
10890        &self,
10891        hidden: &[f32],
10892        position: usize,
10893        logits_out: &mut Vec<f32>,
10894        layers_run: &mut usize,
10895    ) -> Option<Result<Vec<f32>, ()>> {
10896        self.try_token_graph_wgpu_steps(
10897            hidden,
10898            position,
10899            logits_out,
10900            1,
10901            None,
10902            Some(layers_run),
10903            0,
10904            self.num_layers,
10905        )
10906    }
10907
10908    /// The span twin (network split): the graph covers [from..upto_excl)
10909    /// — one submit per SEGMENT per token. lm_head folds in only when
10910    /// the span reaches the last layer.
10911    fn try_token_graph_wgpu_span(
10912        &self,
10913        hidden: &[f32],
10914        position: usize,
10915        logits_out: &mut Vec<f32>,
10916        from: usize,
10917        upto_excl: usize,
10918        layers_run: &mut usize,
10919    ) -> Option<Result<Vec<f32>, ()>> {
10920        self.try_token_graph_wgpu_steps(
10921            hidden,
10922            position,
10923            logits_out,
10924            1,
10925            None,
10926            Some(layers_run),
10927            from,
10928            upto_excl,
10929        )
10930    }
10931
10932    /// Greedy burst: forward `t_next` and let the device pick + re-embed
10933    /// the next k−1 tokens — k frames, ONE submit, k ids back. The ZML
10934    /// trade, on wgpu. None ⇒ caller keeps the per-token path.
10935    fn try_multi_burst(&self, t_next: u32, position: usize, k: usize) -> Option<Vec<u32>> {
10936        if self.o1_active() || self.attn_softcap > 0.0 {
10937            return None;
10938        }
10939        // The burst builds the whole-token graph; attention the graph's
10940        // per-layer geometry cannot express keeps the per-token path.
10941        if let Some(reason) = self.wgpu_graph_attn_decline() {
10942            self.note_graph_decline("wgpu multi-burst", reason);
10943            return None;
10944        }
10945        let graph_on = crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode);
10946        if !graph_on || self.graph_refused() {
10947            // Same memo as the decode site: this path builds the very
10948            // same graph, so a model it cannot build for must not be
10949            // walked again here either. Missing this guard was worth
10950            // 2.5x on an Adreno — 0.361 tok/s against 0.905 — because
10951            // the burst retried per token what decode had already given
10952            // up on.
10953            return None;
10954        }
10955        let emb = self.embed_single(t_next);
10956        let mut lg = Vec::new();
10957        let mut ids = Vec::new();
10958        match self.try_token_graph_wgpu_steps(
10959            &emb,
10960            position,
10961            &mut lg,
10962            k,
10963            Some(&mut ids),
10964            None,
10965            0,
10966            self.num_layers,
10967        ) {
10968            Some(Ok(_)) => {}
10969            Some(Err(())) => {
10970                // Preserve the backend's post-admission failure through the
10971                // Option-based burst API.  The decode caller consumes this
10972                // flag and clears the sequence instead of falling through
10973                // to a stale CPU recurrent state.
10974                self.graph_failed
10975                    .store(true, std::sync::atomic::Ordering::Relaxed);
10976                return None;
10977            }
10978            None => return None,
10979        }
10980        (ids.len() == k).then_some(ids)
10981    }
10982
10983    /// Multi-step greedy: k whole frames in ONE submit, argmax and re-embed
10984    /// on the device. `ids_out` receives the k winner ids; the hidden/logits
10985    /// outputs are NOT produced in that mode.
10986    fn try_token_graph_wgpu_steps(
10987        &self,
10988        hidden: &[f32],
10989        position: usize,
10990        logits_out: &mut Vec<f32>,
10991        steps: usize,
10992        ids_out: Option<&mut Vec<u32>>,
10993        layers_run: Option<&mut usize>,
10994        from: usize,
10995        upto_excl: usize,
10996    ) -> Option<Result<Vec<f32>, ()>> {
10997        // The bank has already reserved its VRAM. Never build a second
10998        // all-expert arena across bank-owned layers (including bursts).
10999        let upto_excl = match self.mimo_moe.graph_prefix_end() {
11000            Some(end) if end < upto_excl => {
11001                if steps != 1 || layers_run.is_none() || from >= end {
11002                    return None;
11003                }
11004                end
11005            }
11006            _ => upto_excl,
11007        };
11008        // O(1) Nyström decode runs off the sealed state, not the KV cache the
11009        // graph mirrors — never take the graph while o1 is active.
11010        let o1_gpu = std::env::var("CMF_O1_GPU").as_deref() == Ok("1");
11011        if (self.o1_active() && !o1_gpu) || self.attn_softcap > 0.0 {
11012            // Softcapped scores have no graph kernel yet — CPU owns them.
11013            // o1 rides the graph only behind CMF_O1_GPU=1 while the port
11014            // proves itself; without it the CPU path owns o1 as before.
11015            return None;
11016        }
11017        // Per-layer KV heads, narrow V, sinks and sliding windows ride
11018        // `GraphAttn::Full::geom` (the ATTEND_X kernels). Anything that
11019        // geometry cannot express declines here, by name — before the
11020        // per-layer gate existed a sliding/sink model ran the graph as
11021        // full-context attention, fluent and wrong. The caller memoizes
11022        // the refusal.
11023        if let Some(reason) = self.wgpu_graph_attn_decline() {
11024            self.note_graph_decline("wgpu token graph", reason);
11025            return None;
11026        }
11027        // Per-layer sealed o1 state for the graph. During prefill the
11028        // state is still Collecting -> views are None -> the graph
11029        // refuses below and the CPU prefill records the q trace and
11030        // seals, exactly as the o1 design requires.
11031        let o1_views: Vec<Option<Vec<crate::nystrom::O1DeviceView<'_>>>> = (from..upto_excl)
11032            .map(|li| {
11033                if !o1_gpu {
11034                    return None;
11035                }
11036                self.kv_cache.layers[self.phys_layer(li)].o1_views()
11037            })
11038            .collect();
11039        if self.o1_active() && o1_gpu {
11040            // Any o1 layer not sealed (or degenerate exact-only) keeps the
11041            // whole token on the CPU: half-graph forwards would desync.
11042            let want: usize = (from..upto_excl)
11043                .filter(|li| self.kv_cache.layers[self.phys_layer(*li)].o1.is_some())
11044                .count();
11045            let have = o1_views.iter().filter(|v| v.is_some()).count();
11046            if want == 0 || have != want {
11047                // The silent twin of the gpu-side o1 gates, found the
11048                // same way: a 15x decode drop with an empty log. Views
11049                // stay None until the layer's state SEALS, so `have`
11050                // lagging `want` early in a run is the o1 design working
11051                // — but it must say so, or the next reader spends a
11052                // night proving the kernels innocent.
11053                // On CHANGE, not once: the first decline is the legal
11054                // unsealed prefill, and a once-print buries the state
11055                // that matters — what the count reads AFTER the seal.
11056                use std::sync::atomic::{AtomicUsize, Ordering};
11057                static LAST: AtomicUsize = AtomicUsize::new(usize::MAX);
11058                let code = have * 1000 + want;
11059                if LAST.swap(code, Ordering::Relaxed) != code {
11060                    tracing::warn!(
11061                        "o1 graph: {have} of {want} layers sealed — per-op until all seal"
11062                    );
11063                }
11064                return None;
11065            }
11066        }
11067        let nh = self.num_heads;
11068        let (nkv, hd, rd) = self.layer_geom(0);
11069        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11070        let mut layers = Vec::with_capacity(upto_excl - from);
11071        let mut model = None;
11072        let dbg = std::env::var("CMF_GRAPH_DEBUG").is_ok();
11073        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
11074            if let Some((m, i, kind, rs)) = t
11075                .graph_weight()
11076                .or_else(|| t.graph_weight_descriptor())
11077            {
11078                let name = &m.tensors[i].name;
11079                let prism = if crate::prism::is_inverse_embedding(m, name) {
11080                    crate::gpu::GraphPrismOp::InverseEmbedding
11081                } else if crate::prism::is_forward_weight(m, name) {
11082                    crate::gpu::GraphPrismOp::Forward
11083                } else {
11084                    crate::gpu::GraphPrismOp::None
11085                };
11086                return Some(crate::gpu::GraphW {
11087                    idx: i,
11088                    kind,
11089                    row_scale: rs,
11090                    data: &[],
11091                    prism,
11092                    affine: crate::prism::is_affine_target(m, name),
11093                });
11094            }
11095            // Small unquantized projections (GDN in_proj_a/b) stay f32.
11096            match t.as_f32() {
11097                Some(d) => Some(crate::gpu::GraphW {
11098                    idx: 0,
11099                    kind: 4,
11100                    row_scale: &[],
11101                    data: d,
11102                    prism: crate::gpu::GraphPrismOp::None,
11103                    affine: false,
11104                }),
11105                None => {
11106                    if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
11107                        eprintln!("batch graph: weight has no graph/f32 representation");
11108                    }
11109                    None
11110                }
11111            }
11112        }
11113        for li in from..upto_excl {
11114            let lw = &self.weights.layers[self.phys_layer(li)];
11115            if dbg {
11116                let ak = match &lw.attn {
11117                    AttnKind::Mla(_) => "Mla".into(),
11118                    AttnKind::Full {
11119                        output_gate, bias, ..
11120                    } => format!("Full gate={output_gate} bias={}", bias.is_some()),
11121                    AttnKind::LinearGdn(_) => "LinearGdn".into(),
11122                    AttnKind::Kda(_) => "Kda".into(),
11123                    AttnKind::Linear(_) => "Linear".into(),
11124                    AttnKind::ShortConv(_) => "ShortConv".into(),
11125                    AttnKind::Bounded(_) => "Bounded".into(),
11126                };
11127                let fk = match &lw.ffn {
11128                    FfnKind::Dense(_) => "Dense",
11129                    FfnKind::Moe(_) => "Moe",
11130                    FfnKind::DenseMoe(_) => "DenseMoe",
11131                };
11132                eprintln!("graph L{li}: attn={ak} ffn={fk}");
11133            }
11134            let gffn = match &lw.ffn {
11135                FfnKind::DenseMoe(_) => return None, // dual branch: CPU path
11136                // A tube layer is several matrices, not one — the
11137                // whole-layer graph has no shape for it yet.
11138                FfnKind::Dense(d) if !d.segs.is_empty() => return None,
11139                FfnKind::Dense(d) => {
11140                    // An activation the graph has no kernel arm for keeps
11141                    // the CPU/per-op owner, by name — the dense graph FFN
11142                    // used to compute SiLU for whatever the model asked.
11143                    let Some(act) = d.act.graph_act() else {
11144                        self.note_graph_decline(
11145                            "wgpu token graph",
11146                            "dense FFN activation without a graph kernel",
11147                        );
11148                        return None;
11149                    };
11150                    crate::gpu::GraphFfn::Dense {
11151                        gate: gw(&d.gate_proj)?,
11152                        up: gw(&d.up_proj)?,
11153                        down: gw(&d.down_proj)?,
11154                        act,
11155                    }
11156                }
11157                FfnKind::Moe(m) => {
11158                    // Adaptive τ and expert masks keep the CPU path, where
11159                    // they are implemented. Sigmoid routing with a selection
11160                    // bias (LFM2-MoE / DeepSeek noaux_tc), a routed scale ≠ 1
11161                    // and an UNGATED shared expert (HunYuan hy_v3: ×2.826 on
11162                    // the routed mix, the shared expert at weight 1) are all
11163                    // graphed — before, every such token fell to the per-op
11164                    // path whole (145 submits/token on Hy-MT2-30B-A3B).
11165                    if m.route_tau.is_some() || m.mask.is_some() {
11166                        return None;
11167                    }
11168                    let shared = m.shared.as_ref();
11169                    let has_shared = shared.is_some();
11170                    let shared_gated = matches!(shared, Some((_, Some(_))));
11171                    let sgate = match shared {
11172                        Some((_, Some(sg))) => gw(sg)?,
11173                        // No gate (hy_v3) or no shared expert at all: the
11174                        // router weight stands in so the plumbing stays
11175                        // total; the select kernels pin weight 1 or skip.
11176                        _ => gw(&m.router)?,
11177                    };
11178                    let router = gw(&m.router)?;
11179                    // The resident MoE kernels do not yet carry the
11180                    // descriptor-aware transform through router/shared-gate
11181                    // selection.  Refuse the complete layer instead of
11182                    // scoring with an untransformed Prism plane (the dense
11183                    // path has an explicit FWHT boundary below).
11184                    if router.prism != crate::gpu::GraphPrismOp::None
11185                        || sgate.prism != crate::gpu::GraphPrismOp::None
11186                        || router.affine
11187                        || sgate.affine
11188                    {
11189                        tracing::warn!(
11190                            "resident MoE declined: Prism/affine router or shared gate transform is not implemented"
11191                        );
11192                        return None;
11193                    }
11194                    let inter = m.experts.first()?.gate_proj.rows();
11195                    let mut experts = Vec::with_capacity(m.experts.len() + 1);
11196                    // q4t or q4tp, but not both in one layer — the kernels
11197                    // are picked per layer, not per expert.
11198                    let mut q4tp: Option<bool> = None;
11199                    // The mixed 2-bit profile: q2tp gate/up over a q4tp
11200                    // down. Uniform across the layer, like `q4tp` itself.
11201                    let mut gu_q2: Option<bool> = None;
11202                    for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
11203                        if !matches!(e.act, Act::Silu)
11204                            || e.gate_proj.rows() != inter
11205                            || e.up_proj.rows() != inter
11206                        {
11207                            return None;
11208                        }
11209                        // Expert tensors are packed into one resident buffer
11210                        // and the MoE kernels have no transform slot per
11211                        // expert.  Keep the CPU/per-op owner for Prism or
11212                        // affine experts rather than silently using raw bytes.
11213                        for expert_weight in [&e.gate_proj, &e.up_proj, &e.down_proj] {
11214                            let Some((em, ei, _, _)) = expert_weight
11215                                .graph_weight()
11216                                .or_else(|| expert_weight.graph_weight_descriptor())
11217                            else {
11218                                return None;
11219                            };
11220                            let name = &em.tensors[ei].name;
11221                            if crate::prism::is_forward_weight(em, name)
11222                                || crate::prism::is_inverse_embedding(em, name)
11223                                || crate::prism::is_affine_target(em, name)
11224                            {
11225                                tracing::warn!(
11226                                    "resident MoE declined: expert Prism/affine transform is not implemented"
11227                                );
11228                                return None;
11229                            }
11230                        }
11231                        let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
11232                            Some((mm, gi)) => (
11233                                mm,
11234                                gi,
11235                                e.up_proj.mapped_q4t()?.1,
11236                                e.down_proj.mapped_q4t()?.1,
11237                                false,
11238                                false,
11239                            ),
11240                            None => match e.gate_proj.mapped_q2tp() {
11241                                Some((mm, gi)) => (
11242                                    mm,
11243                                    gi,
11244                                    e.up_proj.mapped_q2tp()?.1,
11245                                    e.down_proj.mapped_q4tp()?.1,
11246                                    true,
11247                                    true,
11248                                ),
11249                                None => {
11250                                    let (mm, gi) = e.gate_proj.mapped_q4tp()?;
11251                                    (
11252                                        mm,
11253                                        gi,
11254                                        e.up_proj.mapped_q4tp()?.1,
11255                                        e.down_proj.mapped_q4tp()?.1,
11256                                        true,
11257                                        false,
11258                                    )
11259                                }
11260                            },
11261                        };
11262                        if *q4tp.get_or_insert(is_p) != is_p || *gu_q2.get_or_insert(is_q2) != is_q2
11263                        {
11264                            // The shared expert rides in the same packed
11265                            // buffer as the routed ones, so a layer that
11266                            // mixes layouts cannot be indexed by one stride.
11267                            // Say so: the symptom is a whole model quietly
11268                            // running its MoE on the CPU.
11269                            tracing::warn!(
11270                                "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."
11271                            );
11272                            return None;
11273                        }
11274                        model.get_or_insert_with(|| mm.clone());
11275                        experts.push((gi, ui, di));
11276                    }
11277                    crate::gpu::GraphFfn::Moe {
11278                        router,
11279                        shared_gate: sgate,
11280                        experts,
11281                        n_exp: m.experts.len(),
11282                        // CMF_TOPK_PROBE: timing probe only — output is WRONG.
11283                        // Fewer experts shrink the MoE arithmetic while the
11284                        // dispatch count stays identical, which is the only
11285                        // clean way to tell a launch-bound decode from a
11286                        // compute-bound one.
11287                        top_k: std::env::var("CMF_TOPK_PROBE")
11288                            .ok()
11289                            .and_then(|v| v.parse::<usize>().ok())
11290                            .filter(|k| *k > 0 && *k <= m.top_k)
11291                            .unwrap_or(m.top_k),
11292                        inter,
11293                        norm_topk: m.norm_topk_prob,
11294                        q4tp: q4tp?,
11295                        gu_q2: gu_q2.unwrap_or(false),
11296                        sigmoid: m.router_sigmoid,
11297                        bias: m.expert_bias.as_deref(),
11298                        has_shared,
11299                        shared_gated,
11300                        route_scale: m.routed_scaling,
11301                    }
11302                }
11303            };
11304            let attn = match &lw.attn {
11305                AttnKind::Full {
11306                    wq,
11307                    wk,
11308                    wv,
11309                    wo,
11310                    q_norm,
11311                    k_norm,
11312                    output_gate,
11313                    softplus_gate,
11314                    bias,
11315                } => {
11316                    if self.attention_heads_per_layer.is_some() {
11317                        return None;
11318                    }
11319                    // A projected output gate rides the graph in one form:
11320                    // Spark-X2.5's head-wise sigmoid (one g_proj row per Q
11321                    // head). Laguna's softplus gate keeps the CPU path.
11322                    let head_gate = match softplus_gate {
11323                        None => None,
11324                        Some((g, true)) if self.proj_gate_sigmoid && !*output_gate => {
11325                            Some(gw(g)?)
11326                        }
11327                        Some(_) => {
11328                            self.note_graph_decline(
11329                                "wgpu token graph",
11330                                "projected softplus / per-element output gate",
11331                            );
11332                            return None;
11333                        }
11334                    };
11335                    let (m, _, _, _) = wq
11336                        .graph_weight()
11337                        .or_else(|| wq.graph_weight_descriptor())?;
11338                    model = Some(m.clone());
11339                    crate::gpu::GraphAttn::Full {
11340                        wq: gw(wq)?,
11341                        wk: gw(wk)?,
11342                        wv: gw(wv)?,
11343                        wo: gw(wo)?,
11344                        q_norm: q_norm.as_deref(),
11345                        k_norm: k_norm.as_deref(),
11346                        late_qk_norm: self.qk_norm_after_rope,
11347                        bias: bias
11348                            .as_ref()
11349                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
11350                        output_gate: *output_gate,
11351                        cpu_k: self.kv_cache.layers[li].k_heads(),
11352                        cpu_v: self.kv_cache.layers[li].v_heads(),
11353                        cpu_base: self.kv_cache.layers[li].base(),
11354                        geom: self.graph_attn_geom(li),
11355                        head_gate,
11356                    }
11357                }
11358                AttnKind::LinearGdn(w) => {
11359                    let cfg = self.gdn_cfg?;
11360                    let (m, _, _, _) = w
11361                        .in_proj_qkv
11362                        .graph_weight()
11363                        .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
11364                    model = Some(m.clone());
11365                    crate::gpu::GraphAttn::Gdn {
11366                        qkv: gw(&w.in_proj_qkv)?,
11367                        z: gw(&w.in_proj_z)?,
11368                        a: gw(&w.in_proj_a)?,
11369                        b: gw(&w.in_proj_b)?,
11370                        out: gw(&w.out_proj)?,
11371                        conv1d: &w.conv1d,
11372                        a_log: &w.a_log,
11373                        dt_bias: &w.dt_bias,
11374                        norm: &w.norm,
11375                        nv: cfg.num_v_heads,
11376                        nk: cfg.num_k_heads,
11377                        dk: cfg.key_head_dim,
11378                        dv: cfg.value_head_dim,
11379                        kk: cfg.conv_kernel,
11380                        cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11381                    }
11382                }
11383                AttnKind::ShortConv(w) => {
11384                    let cfg = self.short_conv_cfg?;
11385                    let (m, _, _, _) = w
11386                        .in_proj
11387                        .graph_weight()
11388                        .or_else(|| w.in_proj.graph_weight_descriptor())?;
11389                    model = Some(m.clone());
11390                    crate::gpu::GraphAttn::ShortConv {
11391                        inp: gw(&w.in_proj)?,
11392                        out: gw(&w.out_proj)?,
11393                        taps: &w.conv,
11394                        kernel: cfg.kernel,
11395                        cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
11396                    }
11397                }
11398                _ => return None,
11399            };
11400            layers.push(crate::gpu::GraphLayer {
11401                input_norm: &lw.input_norm,
11402                attn,
11403                post_norm: &lw.post_norm,
11404                ffn: gffn,
11405            });
11406        }
11407        let model = model?;
11408        // Fold final-norm + lm_head into the graph when this call wants logits
11409        // and the lm_head is a graphable (quantized) weight — the graph then
11410        // reads back logits (into logits_out) instead of the hidden, dropping
11411        // the separate CPU/GPU lm_head op + its sync. Never the f32 fallback:
11412        // an unquantized lm_head is vocab·hidden and must not be uploaded.
11413        let lm_gw = if upto_excl == self.num_layers
11414            && self.graph_want_logits
11415            && std::env::var("CMF_GPU_LMHEAD")
11416                .map(|v| v != "0")
11417                .unwrap_or(true)
11418        {
11419            self.weights
11420                .lm_head
11421                .graph_weight()
11422                .or_else(|| self.weights.lm_head.graph_weight_descriptor())
11423                .map(|(m, i, kind, rs)| {
11424                let name = &m.tensors[i].name;
11425                let prism = if crate::prism::is_inverse_embedding(m, name) {
11426                    crate::gpu::GraphPrismOp::InverseEmbedding
11427                } else if crate::prism::is_forward_weight(m, name) {
11428                    crate::gpu::GraphPrismOp::Forward
11429                } else {
11430                    crate::gpu::GraphPrismOp::None
11431                };
11432                (
11433                    crate::gpu::GraphW {
11434                        idx: i,
11435                        kind,
11436                        row_scale: rs,
11437                        data: &[],
11438                        prism,
11439                        affine: crate::prism::is_affine_target(m, name),
11440                    },
11441                    self.weights.lm_head.rows(),
11442                )
11443            })
11444        } else {
11445            None
11446        };
11447        let lm = lm_gw.as_ref().map(|(gw, rows)| (gw, *rows));
11448        // Multi-step re-embeds the winner on the device.
11449        let emb_gw = if steps > 1 {
11450            self.weights
11451                .embed_tokens
11452                .graph_weight()
11453                .or_else(|| self.weights.embed_tokens.graph_weight_descriptor())
11454                .map(|(m, i, kind, rs)| {
11455                    let name = &m.tensors[i].name;
11456                    let prism = if crate::prism::is_inverse_embedding(m, name) {
11457                        crate::gpu::GraphPrismOp::InverseEmbedding
11458                    } else if crate::prism::is_forward_weight(m, name) {
11459                        crate::gpu::GraphPrismOp::Forward
11460                    } else {
11461                        crate::gpu::GraphPrismOp::None
11462                    };
11463                    (
11464                        crate::gpu::GraphW {
11465                            idx: i,
11466                            kind,
11467                            row_scale: rs,
11468                            data: &[],
11469                            prism,
11470                            affine: crate::prism::is_affine_target(m, name),
11471                        },
11472                        self.weights.embed_tokens.rows(),
11473                        self.embed_multiplier,
11474                    )
11475                })
11476        } else {
11477            None
11478        };
11479
11480        // Loop boundaries: virtual layer indices after which final_norm is
11481        // applied (mid-stack only; the GLOBAL last layer's norm folds into
11482        // lm_head). Span-relative — the executor compares its enumerate
11483        // index. A span ending mid-stack keeps its boundary norm even when
11484        // it is the span's own last layer.
11485        let loop_norm_at: Vec<usize> = if self.loop_final_norm {
11486            (from..upto_excl.min(self.num_layers - 1))
11487                .filter(|&li| (li + 1) % self.physical_layers == 0)
11488                .map(|li| li - from)
11489                .collect()
11490        } else {
11491            Vec::new()
11492        };
11493        let mut h = hidden.to_vec();
11494        // The normal decode path only needs the fused lm-head logits.  A
11495        // CMF_LOGIT_DUMP diagnostic, however, promises a prompt-boundary
11496        // post-stack hidden alongside those logits; request the existing
11497        // second readback only for that explicit probe instead of dumping
11498        // the input copy left in `h` by a folded-head graph.
11499        let dump_hidden = std::env::var_os("CMF_LOGIT_DUMP").is_some();
11500        let outcome = crate::gpu::forward_token_graph(
11501            &model,
11502            self.graph_kv_id,
11503            &layers,
11504            &o1_views,
11505            self.o1_epoch,
11506            &self.inv_freq,
11507            &mut h,
11508            nh,
11509            nkv,
11510            hd,
11511            self.attn_scale,
11512            rd,
11513            self.hidden_size,
11514            self.intermediate_size,
11515            position,
11516            self.kv_cache.max_seq_len,
11517            gemma,
11518            self.rms_eps as f32,
11519            lm,
11520            &self.weights.final_norm,
11521            logits_out,
11522            &loop_norm_at,
11523            steps,
11524            emb_gw.as_ref().map(|(gw, rows, m)| (gw, *rows, *m)),
11525            ids_out,
11526            layers_run,
11527            from,
11528            dump_hidden,
11529        );
11530        match outcome {
11531            crate::gpu::TokenGraphOutcome::Completed => Some(Ok(h)),
11532            crate::gpu::TokenGraphOutcome::Failed => Some(Err(())),
11533            crate::gpu::TokenGraphOutcome::Declined => None,
11534        }
11535    }
11536
11537    /// Batched prefill: k contiguous prompt positions through the whole wgpu
11538    /// graph in ONE submit (projections/FFN as GEMMs). `hiddens` is [k·hidden]
11539    /// in/out (embeddings in, layer output out); KV mirror / GDN state advance.
11540    /// false ⇒ unsupported → caller keeps the per-position graph.
11541    /// The b-row Metal graph plan for the whole model: every layer as a
11542    /// GDN run or a full-attention item, all-or-nothing (a layer outside the
11543    /// graph's contract → None, the caller runs plain). Shared by the
11544    /// speculative verify and the batched prefill.
11545    #[cfg(target_os = "macos")]
11546    #[allow(clippy::type_complexity)]
11547    fn metal_rows_plan(
11548        &self,
11549    ) -> Option<(
11550        Vec<MetalRowsItem<'_>>,
11551        std::sync::Arc<cortiq_core::CmfModel>,
11552        Option<crate::gpu_metal::GdnGpuCfg>,
11553    )> {
11554        use crate::gpu_metal::{AttnGpuLayer, GdnGpuCfg, GdnGpuLayer, MetalFfn};
11555        let graph_force = crate::gpu::q1_force() || crate::gpu::q2tp_gpu_opt_in();
11556        if !graph_force
11557            || !crate::gpu::enabled_here()
11558            || std::env::var("CMF_GPU_BLOCK")
11559                .map(|v| v == "0")
11560                .unwrap_or(false)
11561            || self.attn_softcap > 0.0
11562            || self.o1_active()
11563            || self.swa.is_some()
11564            || self.global_attn.is_some()
11565            || self.attention_heads_per_layer.is_some()
11566            // per-layer KV heads, narrow V, learned sinks (MiMo-V2)
11567            || self.graph_attn_decline_reason().is_some()
11568            || self.attn_v_norm
11569            || self.loop_final_norm
11570        {
11571            return None;
11572        }
11573        let attend_contract = self.head_dim % 4 == 0
11574            && self.head_dim <= 256
11575            && self.rotary_dim >= 2
11576            && self.rotary_dim <= self.head_dim
11577            && (self.rotary_dim / 2) % 32 == 0
11578            && self.num_kv_heads > 0
11579            && self.num_heads % self.num_kv_heads == 0;
11580        if !attend_contract {
11581            return None;
11582        }
11583        let mut plan: Vec<MetalRowsItem> = Vec::new();
11584        let mut model_ref: Option<std::sync::Arc<cortiq_core::CmfModel>> = None;
11585        for li in 0..self.num_layers {
11586            let lw = &self.weights.layers[self.phys_layer(li)];
11587            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
11588                return None;
11589            }
11590            let ffn = match &lw.ffn {
11591                FfnKind::Dense(d) if d.act == Act::Silu && d.segs.is_empty() => {
11592                    let (Some(g), Some(u), Some(dn)) = (
11593                        d.gate_proj.metal_graph_parts(),
11594                        d.up_proj.metal_graph_parts(),
11595                        d.down_proj.metal_graph_parts(),
11596                    ) else {
11597                        return None;
11598                    };
11599                    MetalFfn::Dense {
11600                        gate: g,
11601                        up: u,
11602                        down: dn,
11603                        gelu: false, // SiLU only (the arm above)
11604                    }
11605                }
11606                _ => return None,
11607            };
11608            match &lw.attn {
11609                AttnKind::LinearGdn(w) if self.gdn_cfg.is_some() => {
11610                    let (Some(qkv), Some(z), Some(a), Some(bb), Some(out)) = (
11611                        w.in_proj_qkv.metal_graph_parts(),
11612                        w.in_proj_z.metal_graph_parts(),
11613                        w.in_proj_a.f32_parts(),
11614                        w.in_proj_b.f32_parts(),
11615                        w.out_proj.metal_graph_parts(),
11616                    ) else {
11617                        return None;
11618                    };
11619                    if let QTensor::Mapped { model, .. } = &w.in_proj_qkv {
11620                        model_ref.get_or_insert_with(|| model.clone());
11621                    }
11622                    let gl = GdnGpuLayer {
11623                        attn_norm: &lw.input_norm,
11624                        post_norm: &lw.post_norm,
11625                        qkv,
11626                        z,
11627                        a,
11628                        b: bb,
11629                        out,
11630                        ffn,
11631                        conv1d: &w.conv1d,
11632                        a_log: &w.a_log,
11633                        dt_bias: &w.dt_bias,
11634                        gnorm: &w.norm,
11635                    };
11636                    match plan.last_mut() {
11637                        Some(MetalRowsItem::Gdn { run, .. }) => run.push(gl),
11638                        _ => plan.push(MetalRowsItem::Gdn {
11639                            run: vec![gl],
11640                            first: li,
11641                        }),
11642                    }
11643                }
11644                AttnKind::Full {
11645                    wq,
11646                    wk,
11647                    wv,
11648                    wo,
11649                    q_norm,
11650                    k_norm,
11651                    output_gate,
11652                    softplus_gate: None,
11653                    bias: None,
11654                } => {
11655                    let (Some(pq), Some(pk), Some(pv), Some(po)) =
11656                        (
11657                            wq.metal_graph_parts(),
11658                            wk.metal_graph_parts(),
11659                            wv.metal_graph_parts(),
11660                            wo.metal_graph_parts(),
11661                        )
11662                    else {
11663                        return None;
11664                    };
11665                    if let QTensor::Mapped { model, .. } = wq {
11666                        model_ref.get_or_insert_with(|| model.clone());
11667                    }
11668                    let cache = &self.kv_cache.layers[li];
11669                    if cache.mode != crate::kv_cache::KvMode::F32 || cache.o1.is_some() {
11670                        return None;
11671                    }
11672                    plan.push(MetalRowsItem::Attn {
11673                        l: AttnGpuLayer {
11674                            attn_norm: &lw.input_norm,
11675                            post_norm: &lw.post_norm,
11676                            wq: pq,
11677                            wk: pk,
11678                            wv: pv,
11679                            wo: po,
11680                            ffn,
11681                        },
11682                        li,
11683                        q_norm: q_norm.as_deref(),
11684                        k_norm: k_norm.as_deref(),
11685                        output_gate: *output_gate,
11686                    });
11687                }
11688                _ => return None,
11689            }
11690        }
11691        let model = model_ref?;
11692        let gcfg = self.gdn_cfg.map(|cfg| GdnGpuCfg {
11693            nv: cfg.num_v_heads,
11694            nk: cfg.num_k_heads,
11695            dk: cfg.key_head_dim,
11696            dv: cfg.value_head_dim,
11697            kk: cfg.conv_kernel,
11698            hidden: self.hidden_size,
11699            inter: self.intermediate_size,
11700            c_dim: cfg.conv_dim(),
11701            eps: cfg.rms_eps as f32,
11702            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11703        });
11704        Some((plan, model, gcfg))
11705    }
11706
11707    /// `AttnDeviceParams` for a plan item over the CPU cache as it stands.
11708    #[cfg(target_os = "macos")]
11709    #[allow(clippy::too_many_arguments)]
11710    fn metal_attn_params<'a>(
11711        li: usize,
11712        cache: &'a crate::kv_cache::LayerKvCache,
11713        q_norm: Option<&'a [f32]>,
11714        k_norm: Option<&'a [f32]>,
11715        output_gate: bool,
11716        inv_freq: &'a [f32],
11717        geom: (usize, usize, usize, usize),
11718        pos0: usize,
11719        kv_id: u64,
11720        scale: f32,
11721        eps: f32,
11722        gemma: bool,
11723        late_qk_norm: bool,
11724    ) -> (crate::gpu_metal::AttnDeviceParams<'a>, usize) {
11725        let (nh, nkv, hd, rd) = geom;
11726        let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
11727        let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
11728        let cpu_stored = cpu_k[0].len() / hd;
11729        (
11730            crate::gpu_metal::AttnDeviceParams {
11731                kv_id,
11732                layer: li,
11733                nh,
11734                nkv,
11735                hd,
11736                rd,
11737                position: pos0,
11738                scale,
11739                eps,
11740                gemma,
11741                late_qk_norm,
11742                output_gate,
11743                q_norm,
11744                k_norm,
11745                inv_freq,
11746                cpu_k,
11747                cpu_v,
11748                cpu_stored,
11749                cpu_gen: cache.generation(),
11750                o1: None,
11751                window: None,
11752                head_gate: None,
11753            },
11754            cpu_stored,
11755        )
11756    }
11757
11758    /// Run the rows plan over `hiddens` (b rows at `pos0..`): validate,
11759    /// encode every item, optionally the head, sync. Returns the graph
11760    /// (for the commit / state finish) plus the GDN layer indices and the
11761    /// attention layers with the row count they were encoded against.
11762    #[cfg(target_os = "macos")]
11763    #[allow(clippy::type_complexity)]
11764    fn metal_rows_run(
11765        &mut self,
11766        hiddens: &mut [f32],
11767        pos0: usize,
11768        b: usize,
11769        prefill: bool,
11770        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11771        // Greedy verify: (row length scored, the b argmax ids out) — the
11772        // head's argmax runs on the device and the logits plane is NOT
11773        // read back (`spec.2` stays empty).
11774        mut argmax_out: Option<(usize, &mut Vec<u32>)>,
11775    ) -> MetalRowsRun {
11776        use crate::gpu_metal::{GraphDims, VerifyGraph};
11777        // The previous round's commit may still be replaying into the
11778        // trunk GDN owners on the second queue: this graph reads them
11779        // (zero-copy wraps) and may reallocate them below — collect the
11780        // replay first. Normally already complete (the draft chain ran
11781        // in between); a failed replay is terminal like a failed commit.
11782        if !crate::gpu_metal::wait_replay() {
11783            tracing::error!("Metal rows graph: the pending async replay failed");
11784            return MetalRowsRun::Failed;
11785        }
11786        spec_stamp("v.wait");
11787        // Seed the GDN recurrent records the rows graph reads — GDN layers
11788        // ONLY (the record the CPU path would allocate anyway).  Sizing
11789        // every layer here planted a zero GDN-sized record on a bounded
11790        // anchor of a natively bounded file this path then refused, and
11791        // that record counted as recurrent state on macOS.
11792        let want = self.gdn_cfg.map(|c| c.state_len()).unwrap_or(0);
11793        if want > 0 {
11794            let phys = self.physical_layers.max(1);
11795            for (li, l) in self.kv_cache.layers.iter_mut().enumerate() {
11796                let is_gdn = self
11797                    .weights
11798                    .layers
11799                    .get(li % phys)
11800                    .is_some_and(|lw| matches!(lw.attn, AttnKind::LinearGdn(_)));
11801                if is_gdn && l.linear_state.len() != want {
11802                    l.linear_state = vec![0f32; want];
11803                }
11804            }
11805        }
11806        let Some((plan, model, gcfg)) = self.metal_rows_plan() else {
11807            return MetalRowsRun::Declined;
11808        };
11809        spec_stamp("v.plan");
11810        let dims = GraphDims {
11811            hidden: self.hidden_size,
11812            eps: self.rms_eps as f32,
11813            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
11814        };
11815        let Some(mut graph) = (if prefill {
11816            VerifyGraph::new_prefill(&model, dims, hiddens, b)
11817        } else {
11818            VerifyGraph::new(&model, dims, hiddens, b)
11819        }) else {
11820            return MetalRowsRun::Declined;
11821        };
11822        let geom = (
11823            self.num_heads,
11824            self.num_kv_heads,
11825            self.head_dim,
11826            self.rotary_dim,
11827        );
11828        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
11829        let eps = self.rms_eps as f32;
11830        let kv_id = self.graph_kv_id;
11831        let inv_freq = self.inv_freq.clone();
11832        for item in &plan {
11833            let ok = match item {
11834                MetalRowsItem::Gdn { run, .. } => gcfg
11835                    .as_ref()
11836                    .map(|gc| run.iter().all(|l| graph.gdn_ok(l, gc)))
11837                    .unwrap_or(false),
11838                MetalRowsItem::Attn {
11839                    l,
11840                    li,
11841                    q_norm,
11842                    k_norm,
11843                    output_gate,
11844                } => {
11845                    let (p, _) = Self::metal_attn_params(
11846                        *li,
11847                        &self.kv_cache.layers[*li],
11848                        *q_norm,
11849                        *k_norm,
11850                        *output_gate,
11851                        &inv_freq,
11852                        geom,
11853                        pos0,
11854                        kv_id,
11855                        self.attn_scale,
11856                        eps,
11857                        gemma,
11858                        self.qk_norm_after_rope,
11859                    );
11860                    graph.attn_ok(l, &p)
11861                }
11862            };
11863            if !ok {
11864                use std::sync::atomic::{AtomicBool, Ordering};
11865                static SAID: AtomicBool = AtomicBool::new(false);
11866                if !SAID.swap(true, Ordering::Relaxed) {
11867                    tracing::warn!("metal rows graph: a layer failed preflight — declining");
11868                }
11869                return MetalRowsRun::Declined;
11870            }
11871        }
11872        let lm = match &spec {
11873            Some((lm, _, _)) => {
11874                if !graph.lm_head_ok(*lm) {
11875                    return MetalRowsRun::Declined;
11876                }
11877                Some(*lm)
11878            }
11879            None => None,
11880        };
11881        let mut gdn_layers = Vec::new();
11882        let mut attn_layers = Vec::new();
11883        for item in &plan {
11884            match item {
11885                MetalRowsItem::Gdn { run, first } => {
11886                    let ro: Vec<&[f32]> = self.kv_cache.layers[*first..*first + run.len()]
11887                        .iter()
11888                        .map(|l| l.linear_state.as_slice())
11889                        .collect();
11890                    if !graph.encode_gdn_run_b(run, &ro, gcfg.as_ref().unwrap()) {
11891                        return MetalRowsRun::Declined;
11892                    }
11893                    gdn_layers.extend(*first..*first + run.len());
11894                }
11895                MetalRowsItem::Attn {
11896                    l,
11897                    li,
11898                    q_norm,
11899                    k_norm,
11900                    output_gate,
11901                } => {
11902                    let (p, cpu_stored) = Self::metal_attn_params(
11903                        *li,
11904                        &self.kv_cache.layers[*li],
11905                        *q_norm,
11906                        *k_norm,
11907                        *output_gate,
11908                        &inv_freq,
11909                        geom,
11910                        pos0,
11911                        kv_id,
11912                        self.attn_scale,
11913                        eps,
11914                        gemma,
11915                        self.qk_norm_after_rope,
11916                    );
11917                    if !graph.encode_attn_b(l, &p) {
11918                        return MetalRowsRun::Declined;
11919                    }
11920                    attn_layers.push((*li, cpu_stored));
11921                }
11922            }
11923        }
11924        if let (Some(lm), Some((_, final_norm, _))) = (lm, spec.as_ref()) {
11925            if !graph.encode_lm_head_b(final_norm, lm) {
11926                return MetalRowsRun::Declined;
11927            }
11928            // The device argmax is an OPTIMISATION, never a reason to
11929            // decline the round: if it will not encode, drop it and read
11930            // the logits plane back the old way (the head is encoded
11931            // either way, so the rows are there).
11932            if let Some((n, _)) = argmax_out.as_ref() {
11933                if !graph.encode_argmax_b(*n) {
11934                    argmax_out = None;
11935                }
11936            }
11937        }
11938        spec_stamp("v.enc");
11939        if !graph.sync() {
11940            return MetalRowsRun::Failed;
11941        }
11942        spec_stamp("v.gpu");
11943        match (spec, argmax_out) {
11944            (Some(_), Some((_, ids))) => {
11945                ids.resize(b, 0);
11946                if !graph.read_argmax(ids) {
11947                    return MetalRowsRun::Failed;
11948                }
11949                spec_stamp("v.am");
11950            }
11951            (Some((lm, _, logits)), None) => {
11952                logits.resize(b * lm.1, 0.0);
11953                if !graph.read_logits(logits) {
11954                    return MetalRowsRun::Failed;
11955                }
11956                spec_stamp("v.lg");
11957            }
11958            (None, _) => {}
11959        }
11960        if !graph.read_hidden(hiddens) {
11961            return MetalRowsRun::Failed;
11962        }
11963        spec_stamp("v.hid");
11964        MetalRowsRun::Completed(MetalVerifyPending {
11965            graph,
11966            gdn_layers,
11967            attn_layers,
11968        })
11969    }
11970
11971    /// Native-Metal twin of `try_batch_graph_wgpu`: the b rows through the
11972    /// whole model on the `VerifyGraph` (one submit), the head folded in
11973    /// when `spec` asks; `hiddens` come back as the last layer's output
11974    /// rows, `spec.2` as `[b][lm_rows]` logits. The graph is parked in
11975    /// `metal_verify` for `metal_verify_commit`.
11976    #[cfg(target_os = "macos")]
11977    fn try_batch_graph_metal(
11978        &mut self,
11979        hiddens: &mut [f32],
11980        positions: &[usize],
11981        b: usize,
11982        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
11983        argmax_out: Option<(usize, &mut Vec<u32>)>,
11984    ) -> crate::gpu::BatchGraphOutcome {
11985        let _t0 = std::time::Instant::now();
11986        if positions.len() != b
11987            || positions.windows(2).any(|w| w[1] != w[0] + 1)
11988            || hiddens.len() != b * self.hidden_size
11989        {
11990            return crate::gpu::BatchGraphOutcome::Declined;
11991        }
11992        let pending = match self.metal_rows_run(hiddens, positions[0], b, false, spec, argmax_out) {
11993            MetalRowsRun::Declined => return crate::gpu::BatchGraphOutcome::Declined,
11994            MetalRowsRun::Failed => return crate::gpu::BatchGraphOutcome::Failed,
11995            MetalRowsRun::Completed(pending) => pending,
11996        };
11997        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
11998            eprintln!(
11999                "metal-verify: {:.1} ms | b={b}",
12000                _t0.elapsed().as_secs_f64() * 1e3
12001            );
12002        }
12003        self.metal_verify = Some(pending);
12004        crate::gpu::BatchGraphOutcome::Completed
12005    }
12006
12007    /// Batched prefill on the Metal rows graph: `ids` (≤ 512) at
12008    /// `start_pos..`, states written in place, K/V rows appended to the
12009    /// CPU caches; optional final norm/head logits are returned in `spec`.
12010    /// Declined means no command buffer was admitted; Failed is terminal.
12011    #[cfg(target_os = "macos")]
12012    fn prefill_rows_metal(
12013        &mut self,
12014        ids: &[u32],
12015        start_pos: usize,
12016        spec: Option<((usize, usize, usize), &[f32], &mut Vec<f32>)>,
12017    ) -> MetalPrefillOutcome {
12018        let b = ids.len();
12019        if b == 0 || b > 512 {
12020            return MetalPrefillOutcome::Declined;
12021        }
12022        METAL_PREFILL_CHUNKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12023        let with_head = spec.is_some();
12024        let hs = self.hidden_size;
12025        let mut hiddens = vec![0f32; b * hs];
12026        for (j, &id) in ids.iter().enumerate() {
12027            let e = self.embed_single(id);
12028            hiddens[j * hs..(j + 1) * hs].copy_from_slice(&e);
12029        }
12030        let mut pending = match self.metal_rows_run(&mut hiddens, start_pos, b, true, spec, None) {
12031            MetalRowsRun::Declined => return MetalPrefillOutcome::Declined,
12032            MetalRowsRun::Failed => {
12033                METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12034                return MetalPrefillOutcome::Failed;
12035            }
12036            MetalRowsRun::Completed(pending) => pending,
12037        };
12038        // states are final: copy them to the owners
12039        let idxs = pending.gdn_layers.clone();
12040        let mut outs: Vec<&mut [f32]> = self
12041            .kv_cache
12042            .layers
12043            .iter_mut()
12044            .enumerate()
12045            .filter(|(i, _)| idxs.binary_search(i).is_ok())
12046            .map(|(_, l)| l.linear_state.as_mut_slice())
12047            .collect();
12048        if !pending.graph.finish_states(&mut outs) {
12049            METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12050            return MetalPrefillOutcome::Failed;
12051        }
12052        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12053        // Read every layer before mutating any CPU cache.  A missing mirror
12054        // row is a terminal graph failure, not a reason to append a partial
12055        // prefix and replay the remainder serially.
12056        let mut rows = Vec::with_capacity(pending.attn_layers.len());
12057        for (li, cpu_stored) in &pending.attn_layers {
12058            let mut kbuf = vec![0f32; b * nkv * hd];
12059            let mut vbuf = vec![0f32; b * nkv * hd];
12060            if !crate::gpu_metal::kv_mirror_read_rows(
12061                self.graph_kv_id,
12062                *li,
12063                nkv,
12064                hd,
12065                *cpu_stored,
12066                b,
12067                &mut kbuf,
12068                &mut vbuf,
12069            ) {
12070                METAL_PREFILL_ERRORS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
12071                return MetalPrefillOutcome::Failed;
12072            }
12073            rows.push((*li, *cpu_stored, kbuf, vbuf));
12074        }
12075        for (li, cpu_stored, kbuf, vbuf) in rows {
12076            let cache = &mut self.kv_cache.layers[li];
12077            for r in 0..b {
12078                cache.append(
12079                    &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12080                    &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12081                    &[],
12082                );
12083            }
12084            crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + b);
12085        }
12086        METAL_PREFILL_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
12087        if with_head {
12088            METAL_PREFILL_HEAD_ROWS.fetch_add(b as u64, std::sync::atomic::Ordering::Relaxed);
12089        }
12090        MetalPrefillOutcome::Completed(hiddens)
12091    }
12092
12093    #[cfg(target_os = "macos")]
12094    fn prefill_batch_metal(&mut self, ids: &[u32], start_pos: usize) -> MetalPrefillOutcome {
12095        self.prefill_rows_metal(ids, start_pos, None)
12096    }
12097
12098    /// Exact teacher-forced NLL through the ordinary Metal rows graph.  This
12099    /// is intentionally separate from the serial TokenGraph scorer: every
12100    /// chunk owns a real b-row graph/head completion and the recurrent/KV
12101    /// handoff is committed before the next chunk begins.
12102    #[cfg(target_os = "macos")]
12103    fn nll_batch_metal(&mut self, ids: &[u32], start: usize) -> MetalBatchNllOutcome {
12104        if ids.len() < 2 || self.o1_active() || self.head_clusters.is_some() {
12105            return MetalBatchNllOutcome::Declined;
12106        }
12107        let Some(lm) = self.weights.lm_head.metal_graph_parts() else {
12108            return MetalBatchNllOutcome::Declined;
12109        };
12110        let chunk = std::env::var("CMF_METAL_PREFILL_CHUNK")
12111            .ok()
12112            .and_then(|v| v.parse::<usize>().ok())
12113            .filter(|&v| (1..=512).contains(&v))
12114            .unwrap_or(32);
12115        let final_norm = self.weights.final_norm.clone();
12116        let mut nll = 0.0f64;
12117        let mut count = 0usize;
12118        let mut pos = 0usize;
12119        let mut completed = 0usize;
12120        while pos < ids.len() {
12121            let end = (pos + chunk).min(ids.len());
12122            let mut logits = Vec::new();
12123            let outcome = self.prefill_rows_metal(
12124                &ids[pos..end],
12125                pos,
12126                Some((lm, &final_norm, &mut logits)),
12127            );
12128            match outcome {
12129                MetalPrefillOutcome::Declined => {
12130                    return if completed == 0 {
12131                        MetalBatchNllOutcome::Declined
12132                    } else {
12133                        MetalBatchNllOutcome::Failed(format!(
12134                            "ordinary Metal NLL batch declined after {completed} chunks"
12135                        ))
12136                    };
12137                }
12138                MetalPrefillOutcome::Failed => {
12139                    return MetalBatchNllOutcome::Failed(
12140                        "ordinary Metal NLL batch failed after admission".to_string(),
12141                    );
12142                }
12143                MetalPrefillOutcome::Completed(_) => {}
12144            }
12145            completed += 1;
12146            let vocab = self.vocab_size.min(lm.1);
12147            if logits.len() != (end - pos) * lm.1 || vocab == 0 {
12148                return MetalBatchNllOutcome::Failed(
12149                    "ordinary Metal NLL head returned an invalid shape".to_string(),
12150                );
12151            }
12152            for row in 0..(end - pos) {
12153                let absolute = pos + row;
12154                if absolute < start || absolute + 1 >= ids.len() {
12155                    continue;
12156                }
12157                let lg = &mut logits[row * lm.1..row * lm.1 + vocab];
12158                if let Some(mu) = self.logit_multiplier {
12159                    for v in lg.iter_mut() {
12160                        *v *= mu;
12161                    }
12162                }
12163                if let Some(c) = self.final_softcap {
12164                    for v in lg.iter_mut() {
12165                        *v = c * (*v / c).tanh();
12166                    }
12167                }
12168                let target = ids[absolute + 1] as usize;
12169                if target >= vocab {
12170                    return MetalBatchNllOutcome::Failed(format!(
12171                        "target token {target} exceeds Metal head rows {vocab}"
12172                    ));
12173                }
12174                let max = lg.iter().fold(f32::NEG_INFINITY, |m, &v| m.max(v));
12175                let lse: f64 = lg
12176                    .iter()
12177                    .map(|&v| ((v - max) as f64).exp())
12178                    .sum::<f64>()
12179                    .ln()
12180                    + max as f64;
12181                nll += lse - lg[target] as f64;
12182                count += 1;
12183            }
12184            pos = end;
12185        }
12186        MetalBatchNllOutcome::Completed(nll, count)
12187    }
12188
12189    /// Commit a Metal verify round: replay the GDN recurrences over the
12190    /// `a + 1` accepted positions into the CPU states, append the accepted
12191    /// K/V rows from the mirrors to the CPU caches, re-point the mirrors.
12192    #[cfg(target_os = "macos")]
12193    fn metal_verify_commit(&mut self, a: usize) -> bool {
12194        let Some(mut pending) = self.metal_verify.take() else {
12195            return false;
12196        };
12197        let n = a + 1;
12198        // encode order == ascending layer order (the plan walks 0..layers)
12199        let idxs = pending.gdn_layers.clone();
12200        let mut outs: Vec<&mut [f32]> = self
12201            .kv_cache
12202            .layers
12203            .iter_mut()
12204            .enumerate()
12205            .filter(|(i, _)| idxs.binary_search(i).is_ok())
12206            .map(|(_, l)| l.linear_state.as_mut_slice())
12207            .collect();
12208        if !pending.graph.commit(n, &mut outs) {
12209            return false;
12210        }
12211        spec_stamp("c.replay");
12212        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12213        // Read every layer before mutating any CPU cache.  Missing rows are
12214        // terminal after the replay has executed; never append a partial KV
12215        // prefix and continue on a serial path.
12216        let mut rows = Vec::with_capacity(pending.attn_layers.len());
12217        for (li, cpu_stored) in &pending.attn_layers {
12218            let mut kbuf = vec![0f32; n * nkv * hd];
12219            let mut vbuf = vec![0f32; n * nkv * hd];
12220            if !crate::gpu_metal::kv_mirror_read_rows(
12221                self.graph_kv_id,
12222                *li,
12223                nkv,
12224                hd,
12225                *cpu_stored,
12226                n,
12227                &mut kbuf,
12228                &mut vbuf,
12229            ) {
12230                return false;
12231            }
12232            rows.push((*li, *cpu_stored, kbuf, vbuf));
12233        }
12234        for (li, cpu_stored, kbuf, vbuf) in rows {
12235            let cache = &mut self.kv_cache.layers[li];
12236            for r in 0..n {
12237                cache.append(
12238                    &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12239                    &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12240                    &[],
12241                );
12242            }
12243            crate::gpu_metal::kv_mirror_set_stored(self.graph_kv_id, li, cpu_stored + n);
12244        }
12245        spec_stamp("c.kv");
12246        true
12247    }
12248
12249    /// The round's warm-ups as ONE b-row graph run over the MTP block on
12250    /// Metal: `pairs` = (trunk hidden, next token) at consecutive positions
12251    /// from `first_pos`; the block's input projection is folded in. This
12252    /// half encodes and SUBMITS (no wait); `mtp_warm_batch_finish` waits
12253    /// and pulls the appended K/V rows into the CPU MTP cache. None = the
12254    /// graph declined (nothing submitted, nothing appended).
12255    #[cfg(target_os = "macos")]
12256    fn mtp_warm_batch_submit(
12257        &mut self,
12258        m: &mut MtpModule,
12259        pairs: &[(&[f32], u32)],
12260        first_pos: usize,
12261    ) -> Option<MetalWarmPending> {
12262        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, VerifyGraph};
12263        let b = pairs.len();
12264        if b == 0 || b > 512 || m.kv.mode != crate::kv_cache::KvMode::F32 || m.kv.o1.is_some() {
12265            return None;
12266        }
12267        let AttnKind::Full {
12268            wq,
12269            wk,
12270            wv,
12271            wo,
12272            q_norm,
12273            k_norm,
12274            output_gate,
12275            softplus_gate: None,
12276            bias: None,
12277        } = &m.layer.attn
12278        else {
12279            return None;
12280        };
12281        let FfnKind::Dense(d) = &m.layer.ffn else {
12282            return None;
12283        };
12284        if !d.segs.is_empty() {
12285            return None;
12286        }
12287        let (Some(pq), Some(pk), Some(pv), Some(po)) =
12288            (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12289        else {
12290            return None;
12291        };
12292        let (Some(g), Some(u), Some(dn)) = (
12293            d.gate_proj.q1_parts(),
12294            d.up_proj.q1_parts(),
12295            d.down_proj.q1_parts(),
12296        ) else {
12297            return None;
12298        };
12299        let Some(eh) = m.eh_proj.q1_parts() else {
12300            return None;
12301        };
12302        let QTensor::Mapped { model, .. } = wq else {
12303            return None;
12304        };
12305        let model = model.clone();
12306        let hs = self.hidden_size;
12307        // [enorm(embed(tok)); hnorm(hidden)] rows
12308        let mut cat = vec![0f32; b * 2 * hs];
12309        for (j, (h, tok)) in pairs.iter().enumerate() {
12310            let e = self.embed_single(*tok);
12311            let (ce, ch) = cat[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
12312            inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, ce);
12313            inference::rms_norm_into(h, &m.hnorm, self.rms_eps, self.norm_style, ch);
12314        }
12315        let dims = GraphDims {
12316            hidden: hs,
12317            eps: self.rms_eps as f32,
12318            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12319        };
12320        spec_stamp("w.cat");
12321        let Some(mut graph) = VerifyGraph::new_via_proj(&model, dims, eh, &cat, b) else {
12322            return None;
12323        };
12324        spec_stamp("w.new");
12325        let l = AttnGpuLayer {
12326            attn_norm: &m.layer.input_norm,
12327            post_norm: &m.layer.post_norm,
12328            wq: pq,
12329            wk: pk,
12330            wv: pv,
12331            wo: po,
12332            ffn: MetalFfn::Dense {
12333                gate: g,
12334                up: u,
12335                down: dn,
12336                gelu: d.act == Act::Gelu,
12337            },
12338        };
12339        let (nh, nkv, hd, rd) = (
12340            self.num_heads,
12341            self.num_kv_heads,
12342            self.head_dim,
12343            self.rotary_dim,
12344        );
12345        let inv_freq = self.inv_freq.clone();
12346        let cpu_stored;
12347        {
12348            let cache = &m.kv;
12349            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12350            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12351            cpu_stored = cpu_k[0].len() / hd;
12352            // The cache may LAG the position (rows nobody warmed): the
12353            // pairs land at cpu_stored.. with their true RoPE positions
12354            // first_pos.., exactly what the one-by-one warm does. A cache
12355            // AHEAD of the position is a real inconsistency.
12356            if cpu_stored > first_pos {
12357                spec_stamp("w.decl");
12358                return None;
12359            }
12360            let p = AttnDeviceParams {
12361                kv_id: self.mtp_kv_id(),
12362                layer: Self::MTP_LAYER_BASE,
12363                nh,
12364                nkv,
12365                hd,
12366                rd,
12367                position: first_pos,
12368                scale: self.attn_scale,
12369                eps: self.rms_eps as f32,
12370                gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12371                late_qk_norm: self.qk_norm_after_rope,
12372                output_gate: *output_gate,
12373                q_norm: q_norm.as_deref(),
12374                k_norm: k_norm.as_deref(),
12375                inv_freq: &inv_freq,
12376                cpu_k,
12377                cpu_v,
12378                cpu_stored,
12379                cpu_gen: cache.generation(),
12380                o1: None,
12381                window: None,
12382                head_gate: None,
12383            };
12384            if !graph.attn_ok(&l, &p) || !graph.encode_attn_b(&l, &p) {
12385                return None;
12386            }
12387        }
12388        spec_stamp("w.enc");
12389        if !graph.submit() {
12390            return None;
12391        }
12392        spec_stamp("w.sub");
12393        Some(MetalWarmPending {
12394            graph,
12395            cpu_stored,
12396            b,
12397        })
12398    }
12399
12400    /// Submit and finish in one call (the prefill's MTP warm-up, where
12401    /// nothing runs in between).
12402    #[cfg(target_os = "macos")]
12403    fn mtp_warm_batch_metal(
12404        &mut self,
12405        m: &mut MtpModule,
12406        pairs: &[(&[f32], u32)],
12407        first_pos: usize,
12408    ) -> bool {
12409        match self.mtp_warm_batch_submit(m, pairs, first_pos) {
12410            Some(p) => self.mtp_warm_batch_finish(m, p),
12411            None => false,
12412        }
12413    }
12414
12415    /// Second half of the batched warm-up: wait for the submitted graph,
12416    /// pull its b appended K/V rows into the CPU MTP cache, re-point the
12417    /// mirror. False = the command buffer failed or the rows are missing
12418    /// (nothing appended; the caller falls back to the one-by-one warm).
12419    #[cfg(target_os = "macos")]
12420    fn mtp_warm_batch_finish(&mut self, m: &mut MtpModule, pending: MetalWarmPending) -> bool {
12421        let MetalWarmPending {
12422            mut graph,
12423            cpu_stored,
12424            b,
12425        } = pending;
12426        let (nkv, hd) = (self.num_kv_heads, self.head_dim);
12427        if !graph.sync() {
12428            return false;
12429        }
12430        spec_stamp("w.gpu");
12431        let mut kbuf = vec![0f32; b * nkv * hd];
12432        let mut vbuf = vec![0f32; b * nkv * hd];
12433        if !crate::gpu_metal::kv_mirror_read_rows(
12434            self.mtp_kv_id(),
12435            Self::MTP_LAYER_BASE,
12436            nkv,
12437            hd,
12438            cpu_stored,
12439            b,
12440            &mut kbuf,
12441            &mut vbuf,
12442        ) {
12443            return false;
12444        }
12445        for r in 0..b {
12446            m.kv.append(
12447                &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12448                &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12449                &[],
12450            );
12451        }
12452        crate::gpu_metal::kv_mirror_set_stored(
12453            self.mtp_kv_id(),
12454            Self::MTP_LAYER_BASE,
12455            cpu_stored + b,
12456        );
12457        spec_stamp("w.kv");
12458        true
12459    }
12460
12461    /// A committed token id from the high table (Cyrillic, CJK and the
12462    /// like sit above 131072 in Qwen's vocabulary; Latin subwords past
12463    /// the 65536 cut are rare enough to lose as rejected drafts) switches
12464    /// the draft to the full head for the next 16 tokens; other ids count
12465    /// down. On an M4 the full 660 MB head costs 5.5 ms a draft step
12466    /// against 1.4 for the shortlist, so the streak is kept short.
12467    pub(crate) fn note_draft_id(&mut self, id: u32) {
12468        let cut = Self::draft_vocab_rows(usize::MAX).max(131_072);
12469        if (id as usize) >= cut {
12470            self.draft_full_streak = 16;
12471        } else {
12472            self.draft_full_streak = self.draft_full_streak.saturating_sub(1);
12473        }
12474    }
12475
12476    /// The draft head's rows for the next step: the shortlist, or the full
12477    /// head while `draft_full_streak` runs.
12478    fn draft_head_rows(&self, head_rows: usize) -> usize {
12479        if self.draft_full_streak > 0 {
12480            head_rows
12481        } else {
12482            Self::draft_vocab_rows(head_rows)
12483        }
12484    }
12485
12486    /// Draft-head shortlist size: `CMF_DRAFT_VOCAB` rows (default 65536,
12487    /// capped at the head; 0 = full head).
12488    fn draft_vocab_rows(head_rows: usize) -> usize {
12489        static N: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
12490        let n = *N.get_or_init(|| {
12491            std::env::var("CMF_DRAFT_VOCAB")
12492                .ok()
12493                .and_then(|v| v.parse().ok())
12494                .unwrap_or(65536)
12495        });
12496        if n == 0 { head_rows } else { n.min(head_rows) }
12497    }
12498
12499    /// One MTP block step on the native Metal token graph: block input on
12500    /// the host, the attention layer + FFN device-resident over the MTP
12501    /// mirror, the head folded in when `want_logits`. The appended K/V row
12502    /// is pulled into the CPU MTP cache (owner of record) after the sync.
12503    #[cfg(target_os = "macos")]
12504    fn mtp_step_metal(
12505        &mut self,
12506        m: &mut MtpModule,
12507        hidden: &[f32],
12508        next_token: u32,
12509        position: usize,
12510        want_logits: bool,
12511    ) -> Option<(Vec<f32>, Vec<f32>)> {
12512        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12513        if std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12514            || !crate::gpu::q1_force()
12515            || !crate::gpu::enabled_here()
12516            || self.attn_softcap > 0.0
12517            || self.attention_heads_per_layer.is_some()
12518            || m.kv.mode != crate::kv_cache::KvMode::F32
12519            || m.kv.o1.is_some()
12520        {
12521            return None;
12522        }
12523        let AttnKind::Full {
12524            wq,
12525            wk,
12526            wv,
12527            wo,
12528            q_norm,
12529            k_norm,
12530            output_gate,
12531            softplus_gate: None,
12532            bias: None,
12533        } = &m.layer.attn
12534        else {
12535            return None;
12536        };
12537        let FfnKind::Dense(d) = &m.layer.ffn else {
12538            return None;
12539        };
12540        if d.act != Act::Silu || !d.segs.is_empty() {
12541            return None;
12542        }
12543        let (pq, pk, pv, po) = (
12544            wq.q1_parts()?,
12545            wk.q1_parts()?,
12546            wv.q1_parts()?,
12547            wo.q1_parts()?,
12548        );
12549        let (g, u, dn) = (
12550            d.gate_proj.q1_parts()?,
12551            d.up_proj.q1_parts()?,
12552            d.down_proj.q1_parts()?,
12553        );
12554        let QTensor::Mapped { model, .. } = wq else {
12555            return None;
12556        };
12557        let model = model.clone();
12558        let lm = if want_logits {
12559            Some(self.weights.lm_head.q1_parts()?)
12560        } else {
12561            None
12562        };
12563        let dims = GraphDims {
12564            hidden: self.hidden_size,
12565            eps: self.rms_eps as f32,
12566            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12567        };
12568        // The block input `eh_proj · [enorm(e); hnorm(h)]` rides in the
12569        // graph (one submit a step); the host per-op matvec if it cannot.
12570        let hs = self.hidden_size;
12571        let mut x = vec![0f32; hs];
12572        let mut graph = TokenGraph::new(&model, dims, &x)?;
12573        let mut folded = false;
12574        if let Some(eh) = m.eh_proj.q1_parts() {
12575            let e = self.embed_single(next_token);
12576            let mut cat = vec![0.0f32; 2 * hs];
12577            let (cat_e, cat_h) = cat.split_at_mut(hs);
12578            inference::rms_norm_into(&e, &m.enorm, self.rms_eps, self.norm_style, cat_e);
12579            inference::rms_norm_into(hidden, &m.hnorm, self.rms_eps, self.norm_style, cat_h);
12580            folded = graph.encode_input_proj(eh, &cat);
12581        }
12582        if !folded {
12583            x = self.mtp_block_input(m, hidden, next_token);
12584            graph = TokenGraph::new(&model, dims, &x)?;
12585        }
12586        spec_stamp("d.in");
12587        let l = AttnGpuLayer {
12588            attn_norm: &m.layer.input_norm,
12589            post_norm: &m.layer.post_norm,
12590            wq: pq,
12591            wk: pk,
12592            wv: pv,
12593            wo: po,
12594            ffn: MetalFfn::Dense {
12595                gate: g,
12596                up: u,
12597                down: dn,
12598                gelu: d.act == Act::Gelu,
12599            },
12600        };
12601        let (nh, nkv, hd, rd) = (
12602            self.num_heads,
12603            self.num_kv_heads,
12604            self.head_dim,
12605            self.rotary_dim,
12606        );
12607        let inv_freq = self.inv_freq.clone();
12608        {
12609            let cache = &m.kv;
12610            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12611            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12612            let cpu_stored = cpu_k[0].len() / hd;
12613            let p = AttnDeviceParams {
12614                kv_id: self.mtp_kv_id(),
12615                layer: Self::MTP_LAYER_BASE,
12616                nh,
12617                nkv,
12618                hd,
12619                rd,
12620                position,
12621                scale: self.attn_scale,
12622                eps: self.rms_eps as f32,
12623                gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12624                late_qk_norm: self.qk_norm_after_rope,
12625                output_gate: *output_gate,
12626                q_norm: q_norm.as_deref(),
12627                k_norm: k_norm.as_deref(),
12628                inv_freq: &inv_freq,
12629                cpu_k,
12630                cpu_v,
12631                cpu_stored,
12632                cpu_gen: cache.generation(),
12633                o1: None,
12634                window: None,
12635                head_gate: None,
12636            };
12637            if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12638                return None;
12639            }
12640        }
12641        // The draft's head over a vocabulary SHORTLIST (the first
12642        // CMF_DRAFT_VOCAB rows — BPE ids run roughly by merge rank, so the
12643        // low ids carry the mass): the verify keeps the full head, so a true
12644        // token past the cut is only a rejected draft, never a wrong token.
12645        // 662 MB a step on Qwen3.8 becomes 170 MB at 65536.
12646        let draft_rows = if let Some(lm) = lm {
12647            self.draft_head_rows(lm.1)
12648        } else {
12649            0
12650        };
12651        if let Some(lm) = lm {
12652            if !graph.lm_head_ok(lm) {
12653                return None;
12654            }
12655            if draft_rows < lm.1 {
12656                if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12657                    return None;
12658                }
12659            } else {
12660                graph.encode_lm_head(&m.final_norm, lm);
12661            }
12662        }
12663        spec_stamp("d.enc");
12664        if graph.sync_checked().is_err() {
12665            return None;
12666        }
12667        spec_stamp("d.gpu");
12668        let mut logits = Vec::new();
12669        if let Some(lm) = lm {
12670            let n_read = draft_rows.min(lm.1).min(self.vocab_size);
12671            logits = attention::take_buf(n_read);
12672            graph.read_logits(&mut logits);
12673            // ids past the shortlist: never drafted (−∞ in every chain)
12674            logits.resize(self.vocab_size, f32::NEG_INFINITY);
12675        }
12676        graph.finish(&mut x);
12677        let mut krow = attention::take_buf(nkv * hd);
12678        let mut vrow = attention::take_buf(nkv * hd);
12679        if crate::gpu_metal::kv_mirror_read_last(
12680            self.mtp_kv_id(),
12681            Self::MTP_LAYER_BASE,
12682            nkv,
12683            hd,
12684            &mut krow,
12685            &mut vrow,
12686        ) {
12687            m.kv.append(&krow, &vrow, &[]);
12688        }
12689        attention::recycle_buf(&mut krow);
12690        attention::recycle_buf(&mut vrow);
12691        spec_stamp("d.rd");
12692        Some((logits, x))
12693    }
12694
12695    /// `CMF_MTP_CHAIN=0` keeps the per-step draft (one submit and one
12696    /// host round trip per MTP step); the default drafts the whole chain
12697    /// in one command buffer when the round is plain greedy.
12698    ///
12699    /// Measured on an M4 (24 GB), Qwen3.8-27B q4tp, P3 at 160 tokens,
12700    /// k=7, six runs per arm alternating inside one lock window — the
12701    /// round's draft phase (median over the 34 rounds of a run) is
12702    /// 34.5 ms per round old against 30.1 new, i.e. 4.93 → 4.31 ms per
12703    /// draft step. That is the whole prize: the 7 submits cost ~0.6 ms
12704    /// each in host and submit latency and nothing else changes —
12705    /// acceptance (3.41 of 7) and tokens per round (4.41) are identical,
12706    /// and the round is 289 → 285 ms, decode 13.8 → 14.0 tok/s.
12707    fn mtp_chain_on() -> bool {
12708        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
12709        *ON.get_or_init(|| std::env::var("CMF_MTP_CHAIN").as_deref() != Ok("0"))
12710    }
12711
12712    /// The round's k greedy drafts as ONE command buffer on Metal: the MTP
12713    /// block k times back to back, each step's token embedding gathered
12714    /// on the device from the argmax the step before it wrote, the head
12715    /// over the round's shortlist (or the full head during a full-head
12716    /// streak — decided once, before the chain, exactly as the per-step
12717    /// path decides it per step, since `draft_full_streak` only moves on
12718    /// a commit). One wait, then the k ids and the k appended K/V rows
12719    /// come back; the CPU MTP cache ends where k `mtp_step_metal` calls
12720    /// would have left it. `Err(false)` = declined before anything was
12721    /// committed (the per-step path takes the round); `Err(true)` = the
12722    /// command buffer failed after commit.
12723    #[cfg(target_os = "macos")]
12724    fn mtp_draft_chain_metal(
12725        &mut self,
12726        m: &mut MtpModule,
12727        hidden: &[f32],
12728        t_next: u32,
12729        position: usize,
12730        k: usize,
12731    ) -> Result<Vec<u32>, bool> {
12732        use crate::gpu_metal::{AttnDeviceParams, AttnGpuLayer, GraphDims, MetalFfn, TokenGraph};
12733        if k == 0
12734            || k > 64
12735            || !Self::mtp_chain_on()
12736            || std::env::var("CMF_MTP_GRAPH").as_deref() == Ok("0")
12737            || !crate::gpu::q1_force()
12738            || !crate::gpu::enabled_here()
12739            || self.attn_softcap > 0.0
12740            || self.attention_heads_per_layer.is_some()
12741            || m.kv.mode != crate::kv_cache::KvMode::F32
12742            || m.kv.o1.is_some()
12743            // the chain gathers embeddings itself: only the plain table
12744            || self.dsv4.is_some()
12745            || self.dsv41.is_some()
12746            || self.qwen4_exp.is_some()
12747            || self.g3n.is_some()
12748        {
12749            return Err(false);
12750        }
12751        let AttnKind::Full {
12752            wq,
12753            wk,
12754            wv,
12755            wo,
12756            q_norm,
12757            k_norm,
12758            output_gate,
12759            softplus_gate: None,
12760            bias: None,
12761        } = &m.layer.attn
12762        else {
12763            return Err(false);
12764        };
12765        let FfnKind::Dense(d) = &m.layer.ffn else {
12766            return Err(false);
12767        };
12768        if d.act != Act::Silu || !d.segs.is_empty() {
12769            return Err(false);
12770        }
12771        let (Some(pq), Some(pk), Some(pv), Some(po)) =
12772            (wq.q1_parts(), wk.q1_parts(), wv.q1_parts(), wo.q1_parts())
12773        else {
12774            return Err(false);
12775        };
12776        let (Some(g), Some(u), Some(dn)) = (
12777            d.gate_proj.q1_parts(),
12778            d.up_proj.q1_parts(),
12779            d.down_proj.q1_parts(),
12780        ) else {
12781            return Err(false);
12782        };
12783        let (Some(eh), Some(lm)) = (m.eh_proj.q1_parts(), self.weights.lm_head.q1_parts()) else {
12784            return Err(false);
12785        };
12786        let QTensor::Mapped { model, .. } = wq else {
12787            return Err(false);
12788        };
12789        let model = model.clone();
12790        // the embedding table: a q4tp tensor of the SAME blob, no Prism
12791        // inverse-embedding post-pass
12792        let QTensor::Mapped {
12793            model: em,
12794            idx: eidx,
12795            dtype: cortiq_core::TensorDtype::Q4TiledP,
12796            ..
12797        } = &self.weights.embed_tokens
12798        else {
12799            return Err(false);
12800        };
12801        if !std::sync::Arc::ptr_eq(em, &model)
12802            || crate::prism::is_inverse_embedding(&model, &model.tensors[*eidx].name)
12803        {
12804            return Err(false);
12805        }
12806        let embed = (
12807            *eidx,
12808            self.weights.embed_tokens.rows(),
12809            self.weights.embed_tokens.cols(),
12810        );
12811        if embed.2 != self.hidden_size || hidden.len() != self.hidden_size {
12812            return Err(false);
12813        }
12814        let dims = GraphDims {
12815            hidden: self.hidden_size,
12816            eps: self.rms_eps as f32,
12817            gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12818        };
12819        let Some(mut graph) = TokenGraph::new(&model, dims, hidden) else {
12820            return Err(false);
12821        };
12822        if !graph.chain_embed_ok(embed) || !graph.lm_head_ok(lm) {
12823            return Err(false);
12824        }
12825        let l = AttnGpuLayer {
12826            attn_norm: &m.layer.input_norm,
12827            post_norm: &m.layer.post_norm,
12828            wq: pq,
12829            wk: pk,
12830            wv: pv,
12831            wo: po,
12832            ffn: MetalFfn::Dense {
12833                gate: g,
12834                up: u,
12835                down: dn,
12836                gelu: d.act == Act::Gelu,
12837            },
12838        };
12839        let (nh, nkv, hd, rd) = (
12840            self.num_heads,
12841            self.num_kv_heads,
12842            self.head_dim,
12843            self.rotary_dim,
12844        );
12845        let inv_freq = self.inv_freq.clone();
12846        let draft_rows = self.draft_head_rows(lm.1);
12847        let n_arg = draft_rows.min(lm.1).min(self.vocab_size);
12848        if n_arg == 0 {
12849            return Err(false);
12850        }
12851        // `CMF_MTP_CHAIN_SPLIT=1` commits each step as it is encoded, so
12852        // the GPU starts on step 0 while the host is still encoding step
12853        // 1 — a probe for whether the host encode is on the critical
12854        // path. It is not: three runs each, draft 30.0 ms per round split
12855        // against 30.1 whole, and the whole chain's host encode measures
12856        // 0.3 ms against a 29.7 ms wait. Kept as a probe, off by default.
12857        let split = std::env::var("CMF_MTP_CHAIN_SPLIT").as_deref() == Ok("1");
12858        let t_chain = std::time::Instant::now();
12859        graph.chain_ids_init(t_next, k);
12860        let cpu_stored;
12861        {
12862            let cache = &m.kv;
12863            let cpu_k: Vec<&[f32]> = (0..nkv).map(|g| cache.head_keys(g)).collect();
12864            let cpu_v: Vec<&[f32]> = (0..nkv).map(|g| cache.head_values(g)).collect();
12865            cpu_stored = cpu_k[0].len() / hd;
12866            for j in 0..k {
12867                if !graph.encode_chain_input(
12868                    embed,
12869                    j as u32,
12870                    &m.enorm,
12871                    &m.hnorm,
12872                    self.embed_multiplier,
12873                    eh,
12874                ) {
12875                    return Err(false);
12876                }
12877                // step j's mirror row: the mirror is re-pointed at the CPU
12878                // rows before step 0 and advances by one per step; its
12879                // resync (never taken past step 0) reads the CPU rows
12880                let p = AttnDeviceParams {
12881                    kv_id: self.mtp_kv_id(),
12882                    layer: Self::MTP_LAYER_BASE,
12883                    nh,
12884                    nkv,
12885                    hd,
12886                    rd,
12887                    position: position + j,
12888                    scale: self.attn_scale,
12889                    eps: self.rms_eps as f32,
12890                    gemma: self.norm_style == cortiq_core::NormStyle::Gemma,
12891                    late_qk_norm: self.qk_norm_after_rope,
12892                    output_gate: *output_gate,
12893                    q_norm: q_norm.as_deref(),
12894                    k_norm: k_norm.as_deref(),
12895                    inv_freq: &inv_freq,
12896                    cpu_k: cpu_k.clone(),
12897                    cpu_v: cpu_v.clone(),
12898                    cpu_stored: cpu_stored + j,
12899                    cpu_gen: cache.generation(),
12900                    o1: None,
12901                    window: None,
12902                    head_gate: None,
12903                };
12904                if !graph.attn_device_ok(&l, &p) || !graph.encode_attn_device(&l, &p) {
12905                    return Err(false);
12906                }
12907                if draft_rows < lm.1 {
12908                    if !graph.encode_lm_head_part(&m.final_norm, lm, draft_rows) {
12909                        return Err(false);
12910                    }
12911                } else {
12912                    graph.encode_lm_head(&m.final_norm, lm);
12913                }
12914                if !graph.encode_argmax(n_arg, j as u32 + 1) {
12915                    return Err(false);
12916                }
12917                if split {
12918                    // CMF_MTP_CHAIN_SPLIT=1: commit every step so the GPU
12919                    // starts on step 0 while the host encodes the rest
12920                    graph.commit();
12921                }
12922            }
12923        }
12924        let t_enc = t_chain.elapsed();
12925        if graph.sync_checked().is_err() {
12926            return Err(true);
12927        }
12928        if std::env::var_os("CMF_GRAPH_SPEC_TIME").is_some() {
12929            eprintln!(
12930                "mtp-chain: encode {:.1} ms | wait {:.1} ms (k={k}, head rows {draft_rows}{})",
12931                t_enc.as_secs_f64() * 1e3,
12932                (t_chain.elapsed() - t_enc).as_secs_f64() * 1e3,
12933                if split { ", split" } else { "" }
12934            );
12935        }
12936        let mut ids = vec![0u32; k];
12937        if !graph.chain_ids_read(&mut ids) {
12938            return Err(true);
12939        }
12940        let mut kbuf = vec![0f32; k * nkv * hd];
12941        let mut vbuf = vec![0f32; k * nkv * hd];
12942        if !crate::gpu_metal::kv_mirror_read_rows(
12943            self.mtp_kv_id(),
12944            Self::MTP_LAYER_BASE,
12945            nkv,
12946            hd,
12947            cpu_stored,
12948            k,
12949            &mut kbuf,
12950            &mut vbuf,
12951        ) {
12952            return Err(true);
12953        }
12954        for r in 0..k {
12955            m.kv.append(
12956                &kbuf[r * nkv * hd..(r + 1) * nkv * hd],
12957                &vbuf[r * nkv * hd..(r + 1) * nkv * hd],
12958                &[],
12959            );
12960        }
12961        Ok(ids)
12962    }
12963
12964    fn try_batch_graph_wgpu(
12965        &self,
12966        hiddens: &mut [f32],
12967        positions: &[usize],
12968        k: usize,
12969        spec: Option<crate::gpu::SpecTail<'_>>,
12970    ) -> crate::gpu::BatchGraphOutcome {
12971        self.try_batch_graph_wgpu_prefix(hiddens, positions, k, spec, None)
12972    }
12973
12974    /// `try_batch_graph_wgpu` with the device-prefix mode: `layers_run`
12975    /// Some lets a stack that does not fit run its leading layers (the
12976    /// token graph's prefix rule) and reports how many; `hiddens` then
12977    /// holds the boundary rows and the caller runs the rest on the host.
12978    fn try_batch_graph_wgpu_prefix(
12979        &self,
12980        hiddens: &mut [f32],
12981        positions: &[usize],
12982        k: usize,
12983        spec: Option<crate::gpu::SpecTail<'_>>,
12984        layers_run: Option<&mut usize>,
12985    ) -> crate::gpu::BatchGraphOutcome {
12986        let graph_end = match self.mimo_moe.graph_prefix_end() {
12987            Some(end) if end < self.num_layers => {
12988                if layers_run.is_none() || spec.is_some() || end == 0 {
12989                    return crate::gpu::BatchGraphOutcome::Declined;
12990                }
12991                end
12992            }
12993            _ => self.num_layers,
12994        };
12995        let _tb = std::time::Instant::now();
12996        let batch_debug = std::env::var_os("CMF_BATCH_DEBUG").is_some();
12997        if self.attn_softcap > 0.0 {
12998            return crate::gpu::BatchGraphOutcome::Declined; // capped scores: no graph kernel — CPU path
12999        }
13000        // Same attention contract as the token graph: per-layer geometry
13001        // rides `geom`, anything it cannot express declines by name.
13002        if let Some(reason) = self.wgpu_graph_attn_decline() {
13003            self.note_graph_decline("wgpu batch graph", reason);
13004            return crate::gpu::BatchGraphOutcome::Declined;
13005        }
13006        let nh = self.num_heads;
13007        let (nkv, hd, rd) = self.layer_geom(0);
13008        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
13009        fn gw(t: &QTensor) -> Option<crate::gpu::GraphW<'_>> {
13010            if let Some((m, i, kind, rs)) = t
13011                .graph_weight()
13012                .or_else(|| t.graph_weight_descriptor())
13013            {
13014                let name = &m.tensors[i].name;
13015                let prism = if crate::prism::is_inverse_embedding(m, name) {
13016                    crate::gpu::GraphPrismOp::InverseEmbedding
13017                } else if crate::prism::is_forward_weight(m, name) {
13018                    crate::gpu::GraphPrismOp::Forward
13019                } else {
13020                    crate::gpu::GraphPrismOp::None
13021                };
13022                return Some(crate::gpu::GraphW {
13023                    idx: i,
13024                    kind,
13025                    row_scale: rs,
13026                    data: &[],
13027                    prism,
13028                    affine: crate::prism::is_affine_target(m, name),
13029                });
13030            }
13031            if std::env::var_os("CMF_BATCH_DEBUG").is_some() {
13032                eprintln!(
13033                    "batch graph: tensor has no graph descriptor/f32 fallback rows={} cols={}",
13034                    t.rows(),
13035                    t.cols()
13036                );
13037            }
13038            t.as_f32().map(|d| crate::gpu::GraphW {
13039                idx: 0,
13040                kind: 4,
13041                row_scale: &[],
13042                data: d,
13043                prism: crate::gpu::GraphPrismOp::None,
13044                affine: false,
13045            })
13046        }
13047        let built: Option<(
13048            Vec<crate::gpu::GraphLayer<'_>>,
13049            std::sync::Arc<cortiq_core::CmfModel>,
13050        )> = (|| {
13051            let mut layers = Vec::with_capacity(graph_end);
13052            let mut model = None;
13053            for li in 0..graph_end {
13054                let lw = &self.weights.layers[self.phys_layer(li)];
13055                // MoE routes per token, so its experts are encoded token by
13056                // token inside the batched submit while attention and the
13057                // projections stay GEMMs. Refusing MoE here is what left
13058                // prefill running one position at a time: 33 tok/s against
13059                // 54 on decode, i.e. reading the prompt was slower than
13060                // writing the answer.
13061                let gffn = match &lw.ffn {
13062                    FfnKind::Dense(d) if !d.segs.is_empty() => {
13063                        if batch_debug {
13064                            eprintln!("batch graph: dense segmented FFN at layer {li}");
13065                        }
13066                        return None;
13067                    }
13068                    FfnKind::Dense(d) => {
13069                        let Some(act) = d.act.graph_act() else {
13070                            if batch_debug {
13071                                eprintln!(
13072                                    "batch graph: dense FFN activation {:?} without a graph kernel at layer {li}",
13073                                    d.act
13074                                );
13075                            }
13076                            return None;
13077                        };
13078                        crate::gpu::GraphFfn::Dense {
13079                            gate: gw(&d.gate_proj)?,
13080                            up: gw(&d.up_proj)?,
13081                            down: gw(&d.down_proj)?,
13082                            act,
13083                        }
13084                    }
13085                    FfnKind::Moe(m) => {
13086                        // Adaptive τ and expert masks stay on the CPU path.
13087                        // Sigmoid scores, the selection bias, a routed scale
13088                        // ≠ 1 and an ungated shared expert (hy_v3) ride the
13089                        // same flags word as the token graph — before, this
13090                        // refusal sent every Hy-MT2-30B prompt to the chunked
13091                        // fallback (8 tok/s of ingest against 53 of decode).
13092                        if m.route_tau.is_some() || m.mask.is_some() {
13093                            return None;
13094                        }
13095                        // A shared expert rides as slot top_k (gated or
13096                        // not is a flag on the select kernel); without one
13097                        // (MiMo-V2, LFM2-MoE) the kernels run top_k slots.
13098                        let shared = m.shared.as_ref();
13099                        let has_shared = shared.is_some();
13100                        let shared_gated = matches!(shared, Some((_, Some(_))));
13101                        let sgate = match shared {
13102                            Some((_, Some(sg))) => gw(sg)?,
13103                            // Ungated or absent: the router plane stands in
13104                            // so the plumbing stays total; the kernel pins
13105                            // weight 1 or never reads it.
13106                            _ => gw(&m.router)?,
13107                        };
13108                        let router = gw(&m.router)?;
13109                        // The batch MoE kernels still consume raw per-token
13110                        // rows and do not carry the descriptor-aware Prism
13111                        // transform/affine bit for router or shared-gate
13112                        // planes.  Refuse rather than route an untransformed
13113                        // source activation.
13114                        if router.prism != crate::gpu::GraphPrismOp::None
13115                            || router.affine
13116                            || sgate.prism != crate::gpu::GraphPrismOp::None
13117                            || sgate.affine
13118                        {
13119                            return None;
13120                        }
13121                        let inter = m.experts.first()?.gate_proj.rows();
13122                        let mut experts = Vec::with_capacity(m.experts.len() + 1);
13123                        let mut q4tp: Option<bool> = None;
13124                        let mut gu_q2: Option<bool> = None;
13125                        for e in m.experts.iter().chain(shared.map(|(se, _)| se)) {
13126                            if !matches!(e.act, Act::Silu)
13127                                || e.gate_proj.rows() != inter
13128                                || e.up_proj.rows() != inter
13129                            {
13130                                return None;
13131                            }
13132                            // Same ladder as the token graph: q4t → q2tp
13133                            // (mixed profile: 2-bit gate/up over a q4tp
13134                            // down) → q4tp. Uniform across the layer.
13135                            let (mm, gi, ui, di, is_p, is_q2) = match e.gate_proj.mapped_q4t() {
13136                                Some((mm, gi)) => (
13137                                    mm,
13138                                    gi,
13139                                    e.up_proj.mapped_q4t()?.1,
13140                                    e.down_proj.mapped_q4t()?.1,
13141                                    false,
13142                                    false,
13143                                ),
13144                                None => match e.gate_proj.mapped_q2tp() {
13145                                    Some((mm, gi)) => (
13146                                        mm,
13147                                        gi,
13148                                        e.up_proj.mapped_q2tp()?.1,
13149                                        e.down_proj.mapped_q4tp()?.1,
13150                                        true,
13151                                        true,
13152                                    ),
13153                                    None => {
13154                                        let (mm, gi) = e.gate_proj.mapped_q4tp()?;
13155                                        (
13156                                            mm,
13157                                            gi,
13158                                            e.up_proj.mapped_q4tp()?.1,
13159                                            e.down_proj.mapped_q4tp()?.1,
13160                                            true,
13161                                            false,
13162                                        )
13163                                    }
13164                                },
13165                            };
13166                            if *q4tp.get_or_insert(is_p) != is_p
13167                                || *gu_q2.get_or_insert(is_q2) != is_q2
13168                            {
13169                                return None;
13170                            }
13171                            if [gi, ui, di].into_iter().any(|idx| {
13172                                mm.tensors
13173                                    .get(idx)
13174                                    .is_some_and(|t| {
13175                                        crate::prism::is_forward_weight(mm, &t.name)
13176                                            || crate::prism::is_affine_target(mm, &t.name)
13177                                    })
13178                            }) {
13179                                return None;
13180                            }
13181                            model.get_or_insert_with(|| mm.clone());
13182                            experts.push((gi, ui, di));
13183                        }
13184                        crate::gpu::GraphFfn::Moe {
13185                            router,
13186                            shared_gate: sgate,
13187                            experts,
13188                            n_exp: m.experts.len(),
13189                            top_k: m.top_k,
13190                            inter,
13191                            norm_topk: m.norm_topk_prob,
13192                            q4tp: q4tp?,
13193                            gu_q2: gu_q2.unwrap_or(false),
13194                            sigmoid: m.router_sigmoid,
13195                            bias: m.expert_bias.as_deref(),
13196                            has_shared,
13197                            shared_gated,
13198                            route_scale: m.routed_scaling,
13199                        }
13200                    }
13201                    _ => return None,
13202                };
13203                let attn = match &lw.attn {
13204                    AttnKind::Full {
13205                        wq,
13206                        wk,
13207                        wv,
13208                        wo,
13209                        q_norm,
13210                        k_norm,
13211                        output_gate,
13212                        softplus_gate,
13213                        bias,
13214                    } => {
13215                        if softplus_gate.is_some() || self.attention_heads_per_layer.is_some() {
13216                            if batch_debug {
13217                                eprintln!(
13218                                    "batch graph: unsupported Full attention gate at layer {li} softplus={} heads={}",
13219                                    softplus_gate.is_some(),
13220                                    self.attention_heads_per_layer.is_some()
13221                                );
13222                            }
13223                            return None;
13224                        }
13225                        let (m, _, _, _) = wq
13226                            .graph_weight()
13227                            .or_else(|| wq.graph_weight_descriptor())?;
13228                        model = Some(m.clone());
13229                        crate::gpu::GraphAttn::Full {
13230                            wq: gw(wq)?,
13231                            wk: gw(wk)?,
13232                            wv: gw(wv)?,
13233                            wo: gw(wo)?,
13234                            q_norm: q_norm.as_deref(),
13235                            k_norm: k_norm.as_deref(),
13236                            late_qk_norm: self.qk_norm_after_rope,
13237                            bias: bias
13238                                .as_ref()
13239                                .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
13240                            output_gate: *output_gate,
13241                            cpu_k: self.kv_cache.layers[li].k_heads(),
13242                            cpu_v: self.kv_cache.layers[li].v_heads(),
13243                            cpu_base: self.kv_cache.layers[li].base(),
13244                            geom: self.graph_attn_geom(li),
13245                            // The batched graph has no head-gate arm; a
13246                            // gated layer is refused above.
13247                            head_gate: None,
13248                        }
13249                    }
13250                    AttnKind::LinearGdn(w) => {
13251                        let Some(cfg) = self.gdn_cfg else {
13252                            if batch_debug {
13253                                eprintln!("batch graph: no GDN config at layer {li}");
13254                            }
13255                            return None;
13256                        };
13257                        let (m, _, _, _) = w
13258                            .in_proj_qkv
13259                            .graph_weight()
13260                            .or_else(|| w.in_proj_qkv.graph_weight_descriptor())?;
13261                        model = Some(m.clone());
13262                        crate::gpu::GraphAttn::Gdn {
13263                            qkv: gw(&w.in_proj_qkv)?,
13264                            z: gw(&w.in_proj_z)?,
13265                            a: gw(&w.in_proj_a)?,
13266                            b: gw(&w.in_proj_b)?,
13267                            out: gw(&w.out_proj)?,
13268                            conv1d: &w.conv1d,
13269                            a_log: &w.a_log,
13270                            dt_bias: &w.dt_bias,
13271                            norm: &w.norm,
13272                            nv: cfg.num_v_heads,
13273                            nk: cfg.num_k_heads,
13274                            dk: cfg.key_head_dim,
13275                            dv: cfg.value_head_dim,
13276                            kk: cfg.conv_kernel,
13277                            cpu_state: &self.kv_cache.layers[self.phys_layer(li)].linear_state,
13278                        }
13279                    }
13280                    _ => return None,
13281                };
13282                layers.push(crate::gpu::GraphLayer {
13283                    input_norm: &lw.input_norm,
13284                    attn,
13285                    post_norm: &lw.post_norm,
13286                    ffn: gffn,
13287                });
13288            }
13289            Some((layers, model?))
13290        })();
13291        let Some((layers, model)) = built else {
13292            {
13293                use std::sync::atomic::{AtomicBool, Ordering};
13294                static SAID: AtomicBool = AtomicBool::new(false);
13295                if !SAID.swap(true, Ordering::Relaxed) {
13296                    tracing::warn!("batch graph: BUILDER refused (layer weights/kinds)");
13297                }
13298            }
13299            return crate::gpu::BatchGraphOutcome::Declined;
13300        };
13301        if std::env::var("CMF_GRAPH_SPEC_TIME").is_ok() {
13302            eprintln!("batch-build: {:.1} ms", _tb.elapsed().as_secs_f64() * 1e3);
13303        }
13304        crate::gpu::forward_batch_graph(
13305            &model,
13306            self.graph_kv_id,
13307            &layers,
13308            &self.inv_freq,
13309            hiddens,
13310            nh,
13311            nkv,
13312            hd,
13313            rd,
13314            self.hidden_size,
13315            self.intermediate_size,
13316            positions,
13317            self.kv_cache.max_seq_len,
13318            gemma,
13319            self.rms_eps as f32,
13320            self.attn_scale,
13321            k,
13322            &(0..graph_end)
13323                .map(|li| self.kv_cache.layers[self.phys_layer(li)].o1_views())
13324                .collect::<Vec<_>>(),
13325            self.o1_epoch,
13326            spec,
13327            layers_run,
13328        )
13329    }
13330
13331    /// Same, stopping after layer `upto` inclusive (routing probe φ).
13332    /// `CMF_DSV4_DRAFT_PROBE=1` — grade the draft against what the trunk goes on
13333    /// to produce. Off by default; it runs a whole draft per decoded token.
13334    fn draft_probe() -> bool {
13335        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13336        *ON.get_or_init(|| std::env::var("CMF_DSV4_DRAFT_PROBE").is_ok_and(|v| v != "0"))
13337    }
13338
13339    /// `CMF_DSV4_DRAFT_PROBE=1`: measure how much of the draft the trunk
13340    /// would have agreed with, WITHOUT verifying or rolling anything back.
13341    ///
13342    /// The number this produces decides the whole speculation design — at
13343    /// acceptance a, a block of B positions yields 1 + a + a² + ... tokens
13344    /// per trunk pass — so it is worth measuring before any of the machinery
13345    /// that would exploit it exists. Each draft is parked with the position
13346    /// it was made at, and graded as the real tokens arrive.
13347    /// `CMF_DSV4_SPEC=1` — the DeepSeek-V4 speculative decode: draft five
13348    /// on the card, verify them in one batched trunk pass, commit the
13349    /// accepted prefix, roll the rest back.
13350    #[cfg(feature = "gpu")]
13351    fn dsv4_spec_on() -> bool {
13352        static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
13353        *ON.get_or_init(|| {
13354            // Test-only runtime gate: model loading still performs the same
13355            // reservation and trunk packing, which gives rollback parity a
13356            // topology-identical non-speculative control arm.
13357            if let Ok(v) = std::env::var("CMF_DSV4_SPEC_RUN") {
13358                return v != "0";
13359            }
13360            // An explicit value is a diagnostic force/escape hatch.  With no
13361            // knob, speculation is eligible only when model loading reserved
13362            // its bounded pack.  On small q4tp cards the geometric reserve
13363            // gate deliberately leaves this at zero: trying to build DSpark
13364            // after the exact trunk filled VRAM is both slower and a device
13365            // OOM (measured on A40).
13366            std::env::var("CMF_DSV4_SPEC")
13367                .map(|v| v != "0")
13368                .unwrap_or_else(|_| {
13369                    crate::gpu_wgpu::DRAFT_RESERVE.load(std::sync::atomic::Ordering::Relaxed) > 0
13370                })
13371        })
13372    }
13373
13374    /// One speculative round at the decode tip. `t_next` is the token the
13375    /// sampler just committed for `next_pos`. Returns the EXTRA accepted
13376    /// tokens (possibly none) and the new position, with `graph_logits`
13377    /// left holding the last accepted position's logits — exactly what the
13378    /// loop top expects. `None` means "speculate not this round": nothing
13379    /// was committed, the caller forwards normally.
13380    #[cfg(feature = "gpu")]
13381    fn dsv4_spec_step(
13382        &mut self,
13383        tip_token: u32,
13384        t_next: u32,
13385        next_pos: usize,
13386        max_extra: usize,
13387        drafted: &mut usize,
13388        accepted_ctr: &mut usize,
13389    ) -> Option<(Vec<u32>, usize)> {
13390        let t_all = std::time::Instant::now();
13391        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13392            thread_local! {
13393                static LAST: std::cell::Cell<Option<std::time::Instant>> =
13394                    const { std::cell::Cell::new(None) };
13395            }
13396            LAST.with(|l| {
13397                if let Some(prev) = l.get() {
13398                    eprintln!(
13399                        "между раундами {:.1} мс",
13400                        prev.elapsed().as_secs_f64() * 1e3
13401                    );
13402                }
13403                l.set(Some(std::time::Instant::now()));
13404            });
13405        }
13406        if std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13407            eprintln!("spec_step: вход pos={next_pos}");
13408        }
13409        let n_layers = self.dsv4.as_ref().map(|b| b.1.len())?;
13410        let cfg = self.dsv4.as_ref().map(|b| b.2)?;
13411        // The draft state and its capture, armed exactly as the probe does.
13412        if self.dspark.is_none() {
13413            let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13414            if t.is_empty() {
13415                return None;
13416            }
13417            crate::dsv4::dspark_arm(&t, cfg.dim);
13418            self.dspark = Some(crate::dsv4::DsparkState::new(
13419                self.dsv4_mtp.len(),
13420                &cfg,
13421                t.len(),
13422            ));
13423        }
13424        let targets = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13425        let pack = crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg);
13426        if pack.is_none() && std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok() {
13427            eprintln!("spec_step: пак не построился (targets {targets:?})");
13428        }
13429        let pack = pack?;
13430        let block = crate::dsv4::dspark_block();
13431        let b_box = self.dsv4.as_mut()?;
13432        let (g, layers, st) = (&b_box.0, &b_box.1, &mut b_box.3);
13433        let ds = self.dspark.as_mut()?;
13434        // The tip's captures: either this token ran on a normal path that
13435        // filled the thread-local, or the previous spec round left them.
13436        let dbg = std::env::var("CMF_DSV4_SPEC_DEBUG").is_ok();
13437        if !crate::dsv4::dspark_take(&mut ds.main_hidden) && !ds.have_hidden {
13438            if dbg {
13439                eprintln!("spec_step: нет захвата");
13440            }
13441            return None;
13442        }
13443        ds.have_hidden = true;
13444        let tip_pos = next_pos.checked_sub(1)?;
13445        let draft_started = std::time::Instant::now();
13446        let mut conf = Vec::new();
13447        let props = crate::dsv4::dspark_draft_gpu(
13448            g,
13449            &self.dsv4_mtp,
13450            &cfg,
13451            ds,
13452            pack,
13453            st.kv_id,
13454            tip_token,
13455            tip_pos,
13456            self.pool.as_deref(),
13457            &mut conf,
13458        );
13459        self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13460        *drafted += block;
13461        if props.is_empty() || props[0] != t_next {
13462            if dbg {
13463                eprintln!(
13464                    "spec_step: черновик {} (props0={:?} t_next={t_next})",
13465                    if props.is_empty() {
13466                        "пуст"
13467                    } else {
13468                        "мимо"
13469                    },
13470                    props.first()
13471                );
13472            }
13473            return None;
13474        }
13475        // `fed[0]` is `t_next`, which the outer loop has already committed;
13476        // only `fed[1..]` become additional output tokens. Cap the verify
13477        // transaction itself to the caller's remaining output budget instead
13478        // of merely truncating the returned vector: otherwise the KV/state
13479        // would advance past `max_tokens` and a 64-token request could return
13480        // 66 tokens (and poison a reused session with two invisible steps).
13481        let mut k_verify = crate::dsv4::dspark_verify_k()
13482            .min(props.len())
13483            .min(max_extra.saturating_add(1));
13484        // Adaptive depth: positions the draft itself doubts are paid for on
13485        // every verify and delivered almost never (natural-text survival
13486        // [.67 .50 .29 .08 .04]). `CMF_DSPARK_CONF_MIN=p` trims the fed
13487        // prefix at the first proposal whose confidence drops below p; on
13488        // predictable text the confidences stay high and nothing changes.
13489        let conf_min = {
13490            static M: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
13491            *M.get_or_init(|| {
13492                std::env::var("CMF_DSPARK_CONF_MIN")
13493                    .ok()
13494                    .and_then(|v| v.parse().ok())
13495                    .unwrap_or(0.0)
13496            })
13497        };
13498        if conf_min > 0.0 && conf.len() >= props.len() {
13499            let mut keep = 1usize;
13500            while keep < k_verify && conf.get(keep).copied().unwrap_or(0.0) >= conf_min {
13501                keep += 1;
13502            }
13503            k_verify = k_verify.min(keep.max(2));
13504        }
13505        if k_verify < 2 {
13506            return None;
13507        }
13508        let mut fed = Vec::with_capacity(k_verify);
13509        fed.push(t_next);
13510        fed.extend_from_slice(&props[1..k_verify]);
13511        let mut argmax = Vec::new();
13512        let mut logits_all = Vec::new();
13513        let mut walked = Vec::new();
13514        let txn = crate::dsv4::dsv4_verify_chunk(
13515            g,
13516            layers,
13517            &cfg,
13518            st,
13519            &fed,
13520            next_pos,
13521            &self.inv_freq,
13522            self.pool.as_deref(),
13523            &targets,
13524            &mut argmax,
13525            &mut logits_all,
13526            &mut walked,
13527        );
13528        if txn.is_none() && dbg {
13529            eprintln!("spec_step: verify отказал");
13530        }
13531        let txn = txn?;
13532        let spec_gpu_end = txn.gpu_end;
13533        let b = fed.len();
13534        let mut accepted = 1usize;
13535        while accepted < b && fed[accepted] == argmax[accepted - 1] {
13536            accepted += 1;
13537        }
13538        // `CMF_DSV4_SPEC_FORCE_REJECT=1` — accept nothing beyond the known
13539        // token, every round: the pure rollback exerciser. The output must
13540        // stay byte-identical to the plain walk; anything else is a
13541        // transaction bug, isolated from the acceptance logic.
13542        if std::env::var("CMF_DSV4_SPEC_FORCE_REJECT").is_ok_and(|v| v != "0") {
13543            accepted = 1;
13544        }
13545        if std::env::var("CMF_DSV4_SPEC_TRACE").is_ok() {
13546            eprintln!("spec@{next_pos}: fed={fed:?} argmax={argmax:?} accepted={accepted}");
13547        }
13548        let t_fin = std::time::Instant::now();
13549        if !crate::dsv4::dsv4_spec_finish(
13550            g,
13551            layers,
13552            &cfg,
13553            st,
13554            txn,
13555            accepted,
13556            &fed,
13557            &self.inv_freq,
13558            self.pool.as_deref(),
13559        ) {
13560            tracing::warn!("dsv4: спекулятивный откат не удался — состояние подозрительно");
13561            return None;
13562        }
13563        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13564            eprintln!(
13565                "finish(k={accepted}): {:.1} мс",
13566                t_fin.elapsed().as_secs_f64() * 1e3
13567            );
13568        }
13569        *accepted_ctr += accepted - 1;
13570        // Captures per accepted token: device targets photographed by the
13571        // batch, host targets from the verify's own walk. The last one
13572        // becomes the new tip's draft input; every one owes the ring an
13573        // entry for its position.
13574        let (hc, dim) = (cfg.hc_mult, cfg.dim);
13575        // Complete-chain layers are photographed by the fused submission;
13576        // partial device layers overwrite that slot after exact host cold-
13577        // expert correction.  Thus every target in the contiguous device
13578        // prefix has a valid per-token capture.
13579        let dev_caps: Vec<usize> = targets
13580            .iter()
13581            .copied()
13582            .filter(|&t| t < spec_gpu_end)
13583            .collect();
13584        let mut caps_all = vec![0.0f32; dev_caps.len() * b * hc * dim];
13585        if !crate::gpu_wgpu::dsv4_spec_cap_read_all(b, dev_caps.len(), hc * dim, &mut caps_all) {
13586            return None;
13587        }
13588        for t in 0..accepted {
13589            let tip = t + 1 == accepted;
13590            for (slot, &tl) in targets.iter().enumerate() {
13591                if let Some(di) = dev_caps.iter().position(|&d| d == tl) {
13592                    let lo = (di * b + t) * hc * dim;
13593                    crate::dsv4::dspark_capture(
13594                        &caps_all[lo..lo + hc * dim],
13595                        &cfg,
13596                        slot,
13597                        &mut ds.main_hidden,
13598                    );
13599                } else if tip
13600                    && crate::dsv4::dspark_peek_slot(slot, dim, {
13601                        let lo = slot * dim;
13602                        &mut ds.main_hidden[lo..lo + dim]
13603                    })
13604                {
13605                    // The tip's host-layer captures are the walk's own
13606                    // per-layer notes — exact. (The walk that ran last ended
13607                    // on exactly this token, on both the accept-all and the
13608                    // rollback path.)
13609                } else {
13610                    // Intermediate tokens: the post-tail state stands in for
13611                    // the per-layer capture on host targets below the last
13612                    // layer. Ring-entry quality only; the tip is exact.
13613                    crate::dsv4::dspark_capture(
13614                        &walked[t * hc * dim..(t + 1) * hc * dim],
13615                        &cfg,
13616                        slot,
13617                        &mut ds.main_hidden,
13618                    );
13619                }
13620            }
13621            crate::dsv4::dspark_ring_append(
13622                g,
13623                &self.dsv4_mtp,
13624                &cfg,
13625                ds,
13626                next_pos + t,
13627                self.pool.as_deref(),
13628            );
13629        }
13630        let row = logits_all[(accepted - 1) * cfg.vocab..accepted * cfg.vocab].to_vec();
13631        self.graph_logits = Some(row);
13632        // The speculative loop never runs the probe, so the trunk tally has
13633        // no other place to cycle. Armed only when someone asked for the
13634        // dump; the host tail is the only tallying path here, which is
13635        // precisely the population a partial pack would serve.
13636        if std::env::var("CMF_DSV4_TRUNK_PICK_DUMP").is_ok() {
13637            crate::dsv4::trunk_freq_note(&crate::dsv4::pick_tally_take());
13638            crate::dsv4::pick_tally_arm();
13639        }
13640        if std::env::var("CMF_DSV4_SPEC_TIME").is_ok() {
13641            eprintln!(
13642                "spec_step total {:.1} мс (k={accepted})",
13643                t_all.elapsed().as_secs_f64() * 1e3
13644            );
13645        }
13646        Some((fed[1..accepted].to_vec(), next_pos + accepted))
13647    }
13648
13649    fn dspark_probe(&mut self, position: usize, token_id: u32) {
13650        if self.dsv4_mtp.is_empty() || !Self::draft_probe() {
13651            return;
13652        }
13653        // What the trunk just routed to, for this token.
13654        let trunk_now = crate::dsv4::pick_tally_take();
13655        crate::dsv4::trunk_freq_note(&trunk_now);
13656        if !trunk_now.is_empty() {
13657            self.dspark_trunk_picks.push(trunk_now);
13658            let keep = crate::dsv4::dspark_block();
13659            if self.dspark_trunk_picks.len() > keep {
13660                self.dspark_trunk_picks.remove(0);
13661            }
13662        }
13663        // Grade whatever is waiting: the token just decoded sits at
13664        // `position`, so it answers the draft made at `position - 1 - i`.
13665        for p in std::mem::take(&mut self.dspark_pending) {
13666            let Some(i) = position.checked_sub(p.0 + 1) else {
13667                continue;
13668            };
13669            let mut p = p;
13670            if i < p.1.len() {
13671                if p.2 && p.1[i] == token_id {
13672                    p.3 = i + 1;
13673                } else {
13674                    p.2 = false;
13675                }
13676                if i + 1 < p.1.len() {
13677                    self.dspark_pending.push(p);
13678                    continue;
13679                }
13680            }
13681            self.dspark_hist.push(p.3);
13682            self.dspark_real.push(token_id);
13683        }
13684        let Some(b) = &mut self.dsv4 else { return };
13685        let (g, layers, cfg) = (&b.0, &b.1, b.2);
13686        let n_layers = layers.len();
13687        if self.dspark.is_none() {
13688            let t = crate::dsv4::dspark_targets(&self.dsv4_mtp, &cfg, n_layers);
13689            if t.is_empty() {
13690                return;
13691            }
13692            eprintln!(
13693                "DSpark: захват со слоёв {t:?}, блок {}",
13694                crate::dsv4::dspark_block()
13695            );
13696            crate::dsv4::dspark_arm(&t, cfg.dim);
13697            self.dspark = Some(crate::dsv4::DsparkState::new(
13698                self.dsv4_mtp.len(),
13699                &cfg,
13700                t.len(),
13701            ));
13702        }
13703        let ds = self.dspark.as_mut().unwrap();
13704        if !crate::dsv4::dspark_take(&mut ds.main_hidden) {
13705            return; // this token ran on a path that captures nothing
13706        }
13707        let mut conf = Vec::new();
13708        crate::dsv4::pick_tally_arm();
13709        // The trunk has already consumed the adaptive VRAM budget. Until the
13710        // draft owns an explicit bounded device pack, its tensors are an
13711        // out-of-core CPU/disk tier by contract: never let per-op probes try
13712        // to squeeze another multi-gigabyte MTP expert cache onto the card.
13713        let draft_started = std::time::Instant::now();
13714        #[cfg(feature = "gpu")]
13715        let gpu_draft = crate::dsv4::dspark_gpu_on();
13716        #[cfg(not(feature = "gpu"))]
13717        let gpu_draft = false;
13718        let props = if gpu_draft {
13719            #[cfg(feature = "gpu")]
13720            {
13721                let kv_id = b.3.kv_id;
13722                match crate::dsv4::dspark_pack_get(&self.dsv4_mtp, &cfg) {
13723                    Some(pk) => crate::dsv4::dspark_draft_gpu(
13724                        g,
13725                        &self.dsv4_mtp,
13726                        &cfg,
13727                        ds,
13728                        pk,
13729                        kv_id,
13730                        token_id,
13731                        position,
13732                        self.pool.as_deref(),
13733                        &mut conf,
13734                    ),
13735                    None => Vec::new(),
13736                }
13737            }
13738            #[cfg(not(feature = "gpu"))]
13739            Vec::new()
13740        } else {
13741            crate::gpu::cpu_scope(|| {
13742                crate::dsv4::dspark_draft(
13743                    g,
13744                    &self.dsv4_mtp,
13745                    &cfg,
13746                    ds,
13747                    token_id,
13748                    position,
13749                    self.pool.as_deref(),
13750                    &mut conf,
13751                )
13752            })
13753        };
13754        self.dspark_draft_ns += draft_started.elapsed().as_nanos();
13755        let draft_picks = crate::dsv4::pick_tally_take();
13756        crate::dsv4::dspark_freq_note(&draft_picks);
13757        // Re-arm for the NEXT trunk token; the probe runs after the forward,
13758        // so this is the only place that can.
13759        crate::dsv4::pick_tally_arm();
13760        if !props.is_empty() {
13761            // Two ratios, side by side: what a batched verify over the trunk
13762            // would read against what it asks for, and the same for the
13763            // draft's three stages. Near 1.0 means a batch amortises nothing.
13764            let (tu, tt) = {
13765                let flat: Vec<(usize, Vec<usize>)> = self
13766                    .dspark_trunk_picks
13767                    .iter()
13768                    .flat_map(|v| v.iter().cloned())
13769                    .collect();
13770                // Per layer, across the window of tokens.
13771                let mut per: std::collections::HashMap<usize, Vec<usize>> =
13772                    std::collections::HashMap::new();
13773                for (li, picks) in flat {
13774                    per.entry(li).or_default().extend(picks);
13775                }
13776                let n = per.len().max(1);
13777                let mut u = 0usize;
13778                let mut t = 0usize;
13779                for (_, v) in per {
13780                    t += v.len();
13781                    u += v.iter().collect::<std::collections::HashSet<_>>().len();
13782                }
13783                (u / n, t / n)
13784            };
13785            let (du, dt) = crate::dsv4::tally_unique(&draft_picks);
13786            self.dspark_exp.push((tu, tt, du, dt));
13787            self.dspark_pending.push((position, props, true, 0));
13788        }
13789        if self.dspark_hist.len() >= 8 && self.dspark_hist.len() % 8 == 0 {
13790            let n = self.dspark_hist.len() as f32;
13791            let mean: f32 = self.dspark_hist.iter().sum::<usize>() as f32 / n;
13792            let block = crate::dsv4::dspark_block();
13793            let mut at = vec![0usize; block + 1];
13794            for &k in &self.dspark_hist {
13795                at[k] += 1;
13796            }
13797            // Prefix survival: S_i = P(the first i positions all held).
13798            let mut surv = Vec::with_capacity(block);
13799            for i in 1..=block {
13800                let k = at[i..].iter().sum::<usize>() as f32 / n;
13801                surv.push(format!("{k:.2}"));
13802            }
13803            let distinct = self
13804                .dspark_real
13805                .iter()
13806                .collect::<std::collections::HashSet<_>>()
13807                .len();
13808            let (tu, tt, du, dt) = self.dspark_exp.iter().fold((0, 0, 0, 0), |a, b| {
13809                (a.0 + b.0, a.1 + b.1, a.2 + b.2, a.3 + b.3)
13810            });
13811            let m = self.dspark_exp.len().max(1);
13812            eprintln!(
13813                "DSpark: черновиков {}, принято в среднем {mean:.2} из {block} \
13814                 (токенов за проход {:.2}), распределение {at:?}, выживание [{}]",
13815                self.dspark_hist.len(),
13816                mean + 1.0,
13817                surv.join(" ")
13818            );
13819            eprintln!(
13820                "DSpark: разных токенов {distinct} из {} (вырожденность), \
13821                 эксперты ствол {}/{} на слой за {block} токенов, \
13822                 черновик {}/{} за блок, draft {:.2} мс/блок",
13823                self.dspark_real.len(),
13824                tu / m,
13825                tt / m,
13826                du / m,
13827                dt / m,
13828                self.dspark_draft_ns as f64 / self.dspark_exp.len().max(1) as f64 / 1e6
13829            );
13830        }
13831    }
13832
13833    fn forward_layers_upto(
13834        &mut self,
13835        hidden: &[f32],
13836        position: usize,
13837        task_mask: Option<&TaskMask>,
13838        upto: Option<usize>,
13839    ) -> Vec<f32> {
13840        // In-process multi-GPU: each segment runs pinned to its card,
13841        // and the only thing crossing the boundary is one hidden vector
13842        // that never leaves this address space. Same layer split the
13843        // network mode does, minus the second process, the socket, the
13844        // serialization and the dir_hash handshake.
13845        if let Some(plan) = self.gpu_plan.clone() {
13846            if upto.is_none() && plan.len() > 1 {
13847                let mut h = hidden.to_vec();
13848                for &(dev, from, upto_incl) in plan.iter() {
13849                    h = crate::gpu::with_device(dev, || {
13850                        self.forward_layers_span(&h, position, task_mask, from, Some(upto_incl))
13851                    });
13852                }
13853                return h;
13854            }
13855        }
13856        self.forward_layers_span(hidden, position, task_mask, 0, upto)
13857    }
13858
13859    /// Split this pipeline's layer stack across local GPUs: segment i
13860    /// runs on `devices[i]`. Contiguous and even by layer count — the
13861    /// VRAM-weighted planner is the next step, and an uneven card pair
13862    /// is why it will be needed. `None` clears the plan.
13863    pub fn set_gpu_plan(&mut self, devices: Option<&[usize]>) -> Result<(), String> {
13864        self.set_gpu_plan_at(devices, None)
13865    }
13866
13867    /// The same, with an explicit first boundary (`--peer-split`): card
13868    /// 0 takes layers `[0..at)`, the rest split what remains. Uneven
13869    /// cards, or an attention-heavy head, are why this knob exists.
13870    pub fn set_gpu_plan_at(
13871        &mut self,
13872        devices: Option<&[usize]>,
13873        at: Option<usize>,
13874    ) -> Result<(), String> {
13875        let Some(devs) = devices.filter(|d| d.len() > 1) else {
13876            self.gpu_plan = None;
13877            return Ok(());
13878        };
13879        self.split_supported()?;
13880        let n = self.num_layers;
13881        if devs.len() > n {
13882            return Err(format!("{} devices for {n} layers", devs.len()));
13883        }
13884        if let Some(k) = at {
13885            if k == 0 || k >= n {
13886                return Err(format!("split at {k}: the model has {n} layers"));
13887            }
13888            if devs.len() == 2 {
13889                self.gpu_plan = Some(std::sync::Arc::new(vec![
13890                    (devs[0], 0, k - 1),
13891                    (devs[1], k, n - 1),
13892                ]));
13893                return Ok(());
13894            }
13895            return Err(format!(
13896                "an explicit split point takes exactly 2 devices, got {}",
13897                devs.len()
13898            ));
13899        }
13900        let per = n.div_ceil(devs.len());
13901        let mut plan = Vec::with_capacity(devs.len());
13902        let mut from = 0usize;
13903        for &d in devs {
13904            if from >= n {
13905                break;
13906            }
13907            let upto = (from + per - 1).min(n - 1);
13908            plan.push((d, from, upto));
13909            from = upto + 1;
13910        }
13911        self.gpu_plan = Some(std::sync::Arc::new(plan));
13912        Ok(())
13913    }
13914
13915    /// The active in-process split, if any: (device, first layer, last).
13916    pub fn gpu_plan(&self) -> Option<Vec<(usize, usize, usize)>> {
13917        self.gpu_plan.as_ref().map(|p| p.as_ref().clone())
13918    }
13919
13920    /// Layer span [from ..= upto] (upto None = last layer): the building
13921    /// block the network pipeline-split rides on. `from > 0` skips the
13922    /// arch escape hatches (the pub `forward_span` refuses those archs
13923    /// first) and the whole-token graph — the plain per-layer loop is
13924    /// the canonical executor for a partial stack.
13925    fn embryo_qtensor_f32(t: &QTensor) -> Vec<f32> {
13926        if let Some(x) = t.as_f32() {
13927            return x.to_vec();
13928        }
13929        let mut out = vec![0.0; t.rows() * t.cols()];
13930        for r in 0..t.rows() {
13931            t.row_f32(r, &mut out[r * t.cols()..(r + 1) * t.cols()]);
13932        }
13933        out
13934    }
13935
13936    fn embryo_resident_eligible(&self) -> bool {
13937        // One mixer family per file: vmf_phase (kind 0/1) or
13938        // gated_delta_net (kind 4); the anchors are full (2) or bounded (3).
13939        if (self.vmf_cfg.is_none() && self.gdn_cfg.is_none())
13940            || self.num_layers != self.physical_layers
13941            || self.loop_final_norm
13942            || self.weights.layers.len() != self.num_layers
13943            || self.head_clusters.is_none()
13944            || self.final_softcap.is_some()
13945            || self.logit_multiplier.is_some()
13946            || self.attn_softcap != 0.0
13947            || self.mtp.is_some()
13948            || self.g3n.is_some()
13949            || self.dsv4.is_some()
13950            || self.dsv41.is_some()
13951            || self.qwen4_exp.is_some()
13952            // Dynamic routing swaps FFN weights mid-sequence under the
13953            // packed graph; a blend has no single overlay. A STATIC skill
13954            // (`from_model_with_skill`) is fine: the pack reads the live
13955            // `weights.layers[*].ffn`, i.e. the skill's tensors, and a
13956            // later `set_active_skill` change drops the pack
13957            // (`invalidate_for_weight_change`).
13958            || self.dyn_router.is_some()
13959            || self.dyn_phi_layer.is_some()
13960            || self.dyn_blend_loaded
13961            || self.o1_cfg.is_some()
13962            || self.swa.is_some()
13963            || self.sliding_layers.is_some()
13964            || self.global_attn.is_some()
13965            || self.attention_heads_per_layer.is_some()
13966            || self.attn_v_norm
13967            || self
13968                .kv_cache
13969                .layers
13970                .iter()
13971                .any(|l| l.mode != crate::kv_cache::KvMode::F32)
13972            || self.rope_scale != 1.0
13973            || self.rope_scale_local != 1.0
13974            || self.attn_scale != 1.0 / (self.head_dim as f32).sqrt()
13975            || self.hidden_size == 0
13976            || self.hidden_size > 1024
13977            || self.intermediate_size > 1024
13978            || self.num_heads == 0
13979            || self.num_kv_heads == 0
13980            || self.num_heads % self.num_kv_heads != 0
13981            || self.num_heads.saturating_mul(self.head_dim) > 1024
13982            || self.num_kv_heads.saturating_mul(self.head_dim) > 1024
13983            || self.vocab_size == 0
13984            || self.kv_cache.max_seq_len == 0
13985            || self.rotary_dim == 0
13986            || self.rotary_dim > self.head_dim
13987            || self.rotary_dim % 2 != 0
13988            || self.inv_freq.len() < self.rotary_dim / 2
13989        {
13990            return false;
13991        }
13992        // The resident shader is deliberately an f32 profile.  Dequantizing
13993        // a Q4/Q8 tensor into the packed buffer would silently change the
13994        // operator relative to the CPU quantized path, so quantized CMFs
13995        // retain the exact ordinary executor instead of claiming parity.
13996        // Measured consequence (RTX PRO 4000, S4 bounded export requantized
13997        // with `cortiq requant --quant q4tp-quantize`): `eligible=false`,
13998        // the generic wgpu whole-token graph refuses too, and the per-op
13999        // path decodes at ~73 tok/s against ~200 tok/s on the CPU q4tp
14000        // path — a q4tp Embryo-O1 file is a CPU artifact today; the
14001        // resident graph serves the f32 export.
14002        if self.weights.lm_head.as_f32().is_none()
14003            || self.weights.embed_tokens.as_f32().is_none()
14004            || self.weights.lm_head.rows() < self.vocab_size
14005            || self.weights.lm_head.cols() != self.hidden_size
14006            || self.weights.embed_tokens.rows() < self.vocab_size
14007            || self.weights.embed_tokens.cols() != self.hidden_size
14008            || self.weights.final_norm.len() != self.hidden_size
14009        {
14010            return false;
14011        }
14012        if let Some(cfg) = self.vmf_cfg {
14013            if cfg.num_heads.saturating_mul(cfg.nphase) > 1024
14014                || cfg.num_heads.saturating_mul(cfg.nphase.saturating_add(1)) > 1024
14015                || cfg.num_heads.saturating_mul(cfg.value_head_dim) > 1024
14016                || cfg.state_len() == 0
14017            {
14018                return false;
14019            }
14020        }
14021        if let Some(g) = self.gdn_cfg {
14022            // The resident GDN kernels (gpu_wgpu.rs `embryo_core_gdn_*`):
14023            // fused projection ≤ 2048 rows, nv·dv ≤ 1024, dk ≤ 128 lanes,
14024            // dv ≤ 256 lanes in vec4 rows, SiLU output gate (the Embryo
14025            // export), same rms eps as the stack.
14026            if g.num_v_heads == 0
14027                || g.num_k_heads == 0
14028                || g.num_v_heads % g.num_k_heads != 0
14029                || g.key_head_dim == 0
14030                || g.key_head_dim > 128
14031                || g.value_head_dim == 0
14032                || g.value_head_dim > 256
14033                || g.value_head_dim % 4 != 0
14034                || g.conv_kernel == 0
14035                || g.num_v_heads > 512
14036                || g.num_v_heads.saturating_mul(g.value_head_dim) > 1024
14037                || g.conv_dim() > 2048
14038                || g.conv_dim() % 4 != 0
14039                || g.hidden_size != self.hidden_size
14040                || g.output_gate_sigmoid
14041                || g.rms_eps != self.rms_eps
14042                || g.state_len() == 0
14043            {
14044                return false;
14045            }
14046        }
14047        let mut full_seen = false;
14048        for lw in &self.weights.layers {
14049            if lw.attn_out_norm.is_some() || lw.ffn_out_norm.is_some() || lw.layer_scale.is_some() {
14050                return false;
14051            }
14052            match &lw.attn {
14053                AttnKind::LinearGdn(w) => {
14054                    let Some(g) = self.gdn_cfg else {
14055                        return false;
14056                    };
14057                    let (nv, dv, kk) = (g.num_v_heads, g.value_head_dim, g.conv_kernel);
14058                    if w.in_proj_qkv.rows() != g.conv_dim()
14059                        || w.in_proj_qkv.cols() != self.hidden_size
14060                        || w.in_proj_qkv.as_f32().is_none()
14061                        || w.in_proj_z.rows() != nv * dv
14062                        || w.in_proj_z.cols() != self.hidden_size
14063                        || w.in_proj_z.as_f32().is_none()
14064                        || w.in_proj_a.rows() != nv
14065                        || w.in_proj_a.cols() != self.hidden_size
14066                        || w.in_proj_a.as_f32().is_none()
14067                        || w.in_proj_b.rows() != nv
14068                        || w.in_proj_b.cols() != self.hidden_size
14069                        || w.in_proj_b.as_f32().is_none()
14070                        || w.conv1d.len() != g.conv_dim() * kk
14071                        || w.a_log.len() != nv
14072                        || w.dt_bias.len() != nv
14073                        || w.norm.len() != dv
14074                        || w.out_proj.rows() != self.hidden_size
14075                        || w.out_proj.cols() != nv * dv
14076                        || w.out_proj.as_f32().is_none()
14077                    {
14078                        return false;
14079                    }
14080                }
14081                AttnKind::Linear(w) => {
14082                    let Some(cfg) = self.vmf_cfg else {
14083                        return false;
14084                    };
14085                    if w.thq.rows() != cfg.num_heads * cfg.nphase
14086                        || w.thq.cols() != self.hidden_size
14087                        || w.thq.as_f32().is_none()
14088                        || w.thk.rows() != cfg.num_heads * cfg.nphase
14089                        || w.thk.cols() != self.hidden_size
14090                        || w.thk.as_f32().is_none()
14091                        || w.v_proj.rows() != cfg.num_heads * cfg.value_head_dim
14092                        || w.v_proj.cols() != self.hidden_size
14093                        || w.v_proj.as_f32().is_none()
14094                        || w.out_proj.rows() != self.hidden_size
14095                        || w.out_proj.cols() != cfg.num_heads * cfg.value_head_dim
14096                        || w.out_proj.as_f32().is_none()
14097                        || w.decay.len() != cfg.num_heads * 2 * cfg.nphase
14098                    {
14099                        return false;
14100                    }
14101                    if let Some((kg, kb)) = &w.k_gate {
14102                        if kg.rows() != cfg.num_heads
14103                            || kg.cols() != self.hidden_size
14104                            || kg.as_f32().is_none()
14105                            || kb.len() != cfg.num_heads
14106                        {
14107                            return false;
14108                        }
14109                    }
14110                    if let Some(conv) = &w.conv {
14111                        if conv.len() % self.hidden_size != 0 || conv.len() / self.hidden_size < 2 {
14112                            return false;
14113                        }
14114                    }
14115                }
14116                AttnKind::Full {
14117                    wq,
14118                    wk,
14119                    wv,
14120                    wo,
14121                    q_norm,
14122                    k_norm,
14123                    output_gate,
14124                    softplus_gate,
14125                    bias,
14126                } => {
14127                    if full_seen
14128                        || q_norm.is_some()
14129                        || k_norm.is_some()
14130                        || *output_gate
14131                        || softplus_gate.is_some()
14132                        || bias.is_some()
14133                        || wq.as_f32().is_none()
14134                        || wk.as_f32().is_none()
14135                        || wv.as_f32().is_none()
14136                        || wo.as_f32().is_none()
14137                        || wq.rows() != self.num_heads * self.head_dim
14138                        || wk.rows() != self.num_kv_heads * self.head_dim
14139                        || wv.rows() != self.num_kv_heads * self.head_dim
14140                        || wq.cols() != self.hidden_size
14141                        || wk.cols() != self.hidden_size
14142                        || wv.cols() != self.hidden_size
14143                        || wo.rows() != self.hidden_size
14144                        || wo.cols() != self.num_heads * self.head_dim
14145                    {
14146                        return false;
14147                    }
14148                    full_seen = true;
14149                }
14150                AttnKind::Bounded(w) => {
14151                    // The resident bounded attend scores S + W lanes in one
14152                    // 256-lane chunk; the format caps S + W at 160.
14153                    let Some(ac) = self.anchor_core.as_ref() else {
14154                        return false;
14155                    };
14156                    if self.bounded_rope.is_none()
14157                        || w.window != ac.window
14158                        || w.sink != ac.sink
14159                        || w.window == 0
14160                        || w.window + w.sink > 256
14161                        || w.sink_k.len() != self.num_kv_heads * w.sink * self.head_dim
14162                        || w.sink_v.len() != self.num_kv_heads * w.sink * self.head_dim
14163                        || w.wq.as_f32().is_none()
14164                        || w.wk.as_f32().is_none()
14165                        || w.wv.as_f32().is_none()
14166                        || w.wo.as_f32().is_none()
14167                        || w.wq.rows() != self.num_heads * self.head_dim
14168                        || w.wk.rows() != self.num_kv_heads * self.head_dim
14169                        || w.wv.rows() != self.num_kv_heads * self.head_dim
14170                        || w.wq.cols() != self.hidden_size
14171                        || w.wk.cols() != self.hidden_size
14172                        || w.wv.cols() != self.hidden_size
14173                        || w.wo.rows() != self.hidden_size
14174                        || w.wo.cols() != self.num_heads * self.head_dim
14175                    {
14176                        return false;
14177                    }
14178                }
14179                _ => return false,
14180            }
14181            match &lw.ffn {
14182                FfnKind::Dense(d) => {
14183                    if d.act != Act::Silu
14184                        || !d.segs.is_empty()
14185                        || d.gate_proj.as_f32().is_none()
14186                        || d.up_proj.as_f32().is_none()
14187                        || d.down_proj.as_f32().is_none()
14188                        || d.gate_proj.rows() != self.intermediate_size
14189                        || d.gate_proj.cols() != self.hidden_size
14190                        || d.up_proj.rows() != self.intermediate_size
14191                        || d.up_proj.cols() != self.hidden_size
14192                        || d.down_proj.rows() != self.hidden_size
14193                        || d.down_proj.cols() != self.intermediate_size
14194                    {
14195                        return false;
14196                    }
14197                }
14198                FfnKind::Moe(m) => {
14199                    if m.resonance.is_none()
14200                        || m.top_k != 1
14201                        || m.router_sigmoid
14202                        || !m.norm_topk_prob
14203                        || m.expert_bias.is_some()
14204                        || m.routed_scaling != 1.0
14205                        || m.route_tau.is_some()
14206                        || m.shared.is_none()
14207                        || m.mask.is_some()
14208                        || m.per_expert_scale.is_some()
14209                        || m.router_input_norm
14210                        || m.experts.is_empty()
14211                        || m.experts.len() > 8
14212                    {
14213                        return false;
14214                    }
14215                    let r = m.resonance.as_ref().unwrap();
14216                    if r.mu.len() != m.experts.len() * self.hidden_size
14217                        || r.bias.len() != m.experts.len()
14218                        || r.u.len() != m.experts.len() * r.k * self.hidden_size
14219                        || r.k > 128
14220                    {
14221                        return false;
14222                    }
14223                    let Some((shared, gate)) = &m.shared else {
14224                        return false;
14225                    };
14226                    if gate.is_some() || shared.act != Act::Silu || !shared.segs.is_empty() {
14227                        return false;
14228                    }
14229                    if shared.gate_proj.as_f32().is_none()
14230                        || shared.up_proj.as_f32().is_none()
14231                        || shared.down_proj.as_f32().is_none()
14232                        || shared.gate_proj.rows() != self.intermediate_size
14233                        || shared.gate_proj.cols() != self.hidden_size
14234                        || shared.up_proj.rows() != self.intermediate_size
14235                        || shared.up_proj.cols() != self.hidden_size
14236                        || shared.down_proj.rows() != self.hidden_size
14237                        || shared.down_proj.cols() != self.intermediate_size
14238                    {
14239                        return false;
14240                    }
14241                    for e in &m.experts {
14242                        if e.act != Act::Silu
14243                            || !e.segs.is_empty()
14244                            || e.gate_proj.as_f32().is_none()
14245                            || e.up_proj.as_f32().is_none()
14246                            || e.down_proj.as_f32().is_none()
14247                            || e.gate_proj.rows() != self.intermediate_size
14248                            || e.gate_proj.cols() != self.hidden_size
14249                            || e.up_proj.rows() != self.intermediate_size
14250                            || e.up_proj.cols() != self.hidden_size
14251                            || e.down_proj.rows() != self.hidden_size
14252                            || e.down_proj.cols() != self.intermediate_size
14253                        {
14254                            return false;
14255                        }
14256                    }
14257                }
14258                FfnKind::DenseMoe(_) => return false,
14259            }
14260            if lw.input_norm.len() != self.hidden_size || lw.post_norm.len() != self.hidden_size {
14261                return false;
14262            }
14263        }
14264        if full_seen && self.anchor_core.is_some() {
14265            return false;
14266        }
14267        full_seen || self.num_layers > 0
14268    }
14269
14270    /// The resident Embryo graph is the owner of this pipeline's forward:
14271    /// the same gate `forward_layers_span` applies before handing a token
14272    /// to `forward_embryo_graph` (both graph phases on, the explicit
14273    /// opt-in, a wgpu device, no earlier refusal, an eligible stack).
14274    fn embryo_resident_wanted(&self) -> bool {
14275        crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14276            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14277            && matches!(
14278                std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14279                Ok("1") | Ok("parallel")
14280            )
14281            && crate::gpu::enabled_here()
14282            && !self.graph_refused()
14283            && self.embryo_resident_eligible()
14284    }
14285
14286    /// Chunked prefill on the resident graph: `ids` from `start` in
14287    /// chunks of `EMBRYO_CHUNK_MAX`, one submit each, projections/FFN as
14288    /// chunk GEMMs and the recurrent layers walked in time on the device.
14289    /// Returns the last position's logits when the whole span ran there.
14290    /// `None` = refused before any device work (the per-position path
14291    /// takes the span).  `CMF_EMBRYO_CHUNK=0` keeps the per-position
14292    /// prefill (A/B and the parity reference).
14293    fn embryo_prefill_chunked(&mut self, ids: &[u32], start: usize) -> Option<Vec<f32>> {
14294        if ids.len() < 2
14295            || std::env::var("CMF_EMBRYO_CHUNK").as_deref() == Ok("0")
14296            || !self.embryo_resident_wanted()
14297        {
14298            return None;
14299        }
14300        let model = self.ensure_embryo_graph()?;
14301        let cmax = std::env::var("CMF_EMBRYO_CHUNK")
14302            .ok()
14303            .and_then(|v| v.parse::<usize>().ok())
14304            .filter(|&v| v >= 1)
14305            .unwrap_or(crate::gpu::EMBRYO_CHUNK_MAX)
14306            .min(crate::gpu::EMBRYO_CHUNK_MAX);
14307        let hs = self.hidden_size;
14308        let n = ids.len();
14309        let mut pos = start;
14310        let mut last = None;
14311        let mut rows = Vec::with_capacity(cmax * hs);
14312        while pos < n {
14313            let end = (pos + cmax).min(n);
14314            rows.clear();
14315            for &id in &ids[pos..end] {
14316                rows.extend_from_slice(&self.embed_single(id));
14317            }
14318            let mut lg = Vec::new();
14319            if !crate::gpu::forward_embryo_graph_chunk(
14320                &model,
14321                self.graph_kv_id,
14322                &rows,
14323                pos,
14324                end - pos,
14325                &mut lg,
14326            ) {
14327                if pos == start {
14328                    if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14329                        eprintln!("embryo-dbg: chunk prefill refused at position {pos}");
14330                    }
14331                    return None;
14332                }
14333                // The device sequence advanced through the earlier chunks;
14334                // a host continuation would mix two owners of the state.
14335                // Fail this sequence alone, leaving no stale state or key.
14336                self.kv_cache.clear();
14337                self.clear_history();
14338                crate::gpu::graph_kv_reset(self.graph_kv_id);
14339                panic!(
14340                    "resident Embryo chunk prefill refused at position {pos}; refusing an in-flight host fallback"
14341                );
14342            }
14343            last = Some(lg);
14344            pos = end;
14345        }
14346        if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14347            eprintln!(
14348                "embryo-dbg: chunk prefill {} ids from {start} in {} submits (chunk {cmax})",
14349                n - start,
14350                (n - start).div_ceil(cmax)
14351            );
14352        }
14353        last
14354    }
14355
14356    fn ensure_embryo_graph(&mut self) -> Option<std::sync::Arc<crate::gpu::EmbryoGraphModel>> {
14357        if self.embryo_graph.is_none() && self.embryo_resident_eligible() {
14358            const UMAX: u32 = u32::MAX;
14359            const HEADER: usize = crate::gpu::EMBRYO_META_HEADER;
14360            const REC: usize = 64;
14361            struct Pack {
14362                data: Vec<f32>,
14363            }
14364            impl Pack {
14365                fn put(&mut self, x: &[f32]) -> u32 {
14366                    if x.is_empty() {
14367                        return u32::MAX;
14368                    }
14369                    let off = self.data.len();
14370                    self.data.extend_from_slice(x);
14371                    off as u32
14372                }
14373            }
14374            // The mixer family of the file: vmf_phase geometry fills the
14375            // phase header words, gated_delta_net fills words 24..29.  A
14376            // file has exactly one linear core, so at most one is live.
14377            let vmf = self.vmf_cfg;
14378            let gdn = self.gdn_cfg;
14379            let mut pack = Pack { data: Vec::new() };
14380            let mut meta = vec![0u32; HEADER];
14381            meta[0] = self.hidden_size as u32;
14382            meta[1] = self.intermediate_size as u32;
14383            meta[2] = self.vocab_size as u32;
14384            meta[3] = self.num_layers as u32;
14385            meta[4] = vmf.map(|c| c.num_heads).unwrap_or(0) as u32;
14386            meta[5] = vmf.map(|c| c.nphase).unwrap_or(0) as u32;
14387            meta[6] = vmf.map(|c| c.value_head_dim).unwrap_or(0) as u32;
14388            if let Some(g) = gdn {
14389                meta[24] = g.num_v_heads as u32;
14390                meta[25] = g.num_k_heads as u32;
14391                meta[26] = g.key_head_dim as u32;
14392                meta[27] = g.value_head_dim as u32;
14393                meta[28] = g.conv_kernel as u32;
14394                meta[29] = g.conv_dim() as u32;
14395            }
14396            meta[7] = self.num_heads as u32;
14397            meta[8] = self.num_kv_heads as u32;
14398            meta[9] = self.head_dim as u32;
14399            meta[10] = self.kv_cache.max_seq_len as u32;
14400            let clusters = self.head_clusters.as_ref().unwrap();
14401            let cluster_count = clusters.len() / self.hidden_size;
14402            if clusters.len() % self.hidden_size != 0
14403                || cluster_count == 0
14404                || cluster_count > 1024
14405                || self.vocab_size % cluster_count != 0
14406                || self.weights.lm_head.rows() < self.vocab_size
14407                || self.weights.final_norm.len() != self.hidden_size
14408            {
14409                return None;
14410            }
14411            meta[11] = cluster_count as u32;
14412            meta[12] = (self.vocab_size / cluster_count) as u32;
14413            meta[13] = vmf.map(|c| c.state_len()).unwrap_or(0) as u32;
14414            meta[16] = self.rotary_dim as u32;
14415            meta[17] = matches!(self.norm_style, NormStyle::Gemma) as u32;
14416            meta[19] = (self.rms_eps as f32).to_bits();
14417            let max_conv = self
14418                .weights
14419                .layers
14420                .iter()
14421                .filter_map(|lw| match &lw.attn {
14422                    AttnKind::Linear(w) => w.conv.as_ref().map(|v| v.len() / self.hidden_size),
14423                    _ => None,
14424                })
14425                .max()
14426                .unwrap_or(1);
14427            // One state slot per recurrent layer: the phase state plus its
14428            // hidden-wide conv ring, or the GDN record `[conv ring | S]`
14429            // (`GdnCfg::state_len`).  A file carries one mixer family, so
14430            // the stride is exactly that family's record and the device
14431            // state buffer equals the header's recurrent bytes.
14432            let phase_stride = vmf
14433                .map(|c| c.state_len() + max_conv.saturating_sub(1) * self.hidden_size)
14434                .unwrap_or(0);
14435            let gdn_stride = gdn.map(|g| g.state_len()).unwrap_or(0);
14436            let state_stride = phase_stride.max(gdn_stride);
14437            // Bounded genome: the KV plane of an anchor is its ring
14438            // `[kvh][W][hd]` K + V, and only anchors own one.  Legacy full
14439            // anchors keep the `max_seq` planes indexed by layer.
14440            let bounded = self.anchor_core.clone();
14441            let (anchor_window, anchor_sink) = bounded
14442                .as_ref()
14443                .map(|ac| (ac.window, ac.sink))
14444                .unwrap_or((0, 0));
14445            let kv_stride = if bounded.is_some() {
14446                2usize
14447                    .saturating_mul(self.num_kv_heads)
14448                    .saturating_mul(anchor_window)
14449                    .saturating_mul(self.head_dim)
14450            } else {
14451                2usize
14452                    .saturating_mul(self.num_kv_heads)
14453                    .saturating_mul(self.kv_cache.max_seq_len)
14454                    .saturating_mul(self.head_dim)
14455            };
14456            meta[14] = state_stride as u32;
14457            meta[15] = kv_stride as u32;
14458            meta[18] = anchor_window as u32;
14459            meta[20] = anchor_sink as u32;
14460            meta[21] = match &self.bounded_rope {
14461                Some(rope) => {
14462                    // [W][half] cos then [W][half] sin, one contiguous table.
14463                    let off = pack.put(&rope.cos);
14464                    let _ = pack.put(&rope.sin);
14465                    off
14466                }
14467                None => UMAX,
14468            };
14469            let mut full_seen = false;
14470            let mut bounded_seen = 0usize;
14471            // Recurrent state slots belong to mixer layers only (phase or
14472            // GDN): an anchor owns a ring, not a state stride, so the
14473            // device state buffer is exactly the header's recurrent bytes.
14474            let mut phase_seen = 0usize;
14475            let mut gdn_seen = 0usize;
14476            for (li, lw) in self.weights.layers.iter().enumerate() {
14477                let base = meta.len();
14478                meta.resize(base + REC, UMAX);
14479                meta[base] = match &lw.attn {
14480                    AttnKind::Linear(w) if w.phase_delta => 1,
14481                    AttnKind::Linear(_) => 0,
14482                    AttnKind::Full { .. } => 2,
14483                    AttnKind::Bounded(_) => 3,
14484                    AttnKind::LinearGdn(_) => 4,
14485                    _ => UMAX,
14486                };
14487                meta[base + 1] = pack.put(&lw.input_norm);
14488                meta[base + 2] = pack.put(&lw.post_norm);
14489                meta[base + 25] = match &lw.attn {
14490                    AttnKind::Linear(_) => {
14491                        let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14492                        phase_seen += 1;
14493                        off
14494                    }
14495                    AttnKind::LinearGdn(_) => {
14496                        let off = ((phase_seen + gdn_seen) * state_stride) as u32;
14497                        gdn_seen += 1;
14498                        off
14499                    }
14500                    _ => UMAX,
14501                };
14502                match &lw.attn {
14503                    AttnKind::LinearGdn(w) => {
14504                        // Layer record words 56..63 + 29, as the resident
14505                        // kernels read them (gpu_wgpu.rs `embryo_core_gdn_*`).
14506                        meta[base + 56] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_qkv));
14507                        meta[base + 57] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_z));
14508                        meta[base + 58] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_a));
14509                        meta[base + 59] = pack.put(&Self::embryo_qtensor_f32(&w.in_proj_b));
14510                        meta[base + 60] = pack.put(&w.conv1d);
14511                        meta[base + 61] = pack.put(&w.a_log);
14512                        meta[base + 62] = pack.put(&w.dt_bias);
14513                        meta[base + 63] = pack.put(&w.norm);
14514                        meta[base + 29] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14515                        meta[base + 24] = 0;
14516                    }
14517                    AttnKind::Linear(w) => {
14518                        meta[base + 3] = pack.put(&Self::embryo_qtensor_f32(&w.thq));
14519                        meta[base + 4] = pack.put(&Self::embryo_qtensor_f32(&w.thk));
14520                        meta[base + 5] = pack.put(&Self::embryo_qtensor_f32(&w.v_proj));
14521                        meta[base + 6] = pack.put(&Self::embryo_qtensor_f32(&w.out_proj));
14522                        let decay: Vec<f32> = w.decay.iter().map(|&x| x as f32).collect();
14523                        meta[base + 7] = pack.put(&decay);
14524                        if let Some((kg, kb)) = &w.k_gate {
14525                            meta[base + 8] = pack.put(&Self::embryo_qtensor_f32(kg));
14526                            meta[base + 9] = pack.put(kb);
14527                        }
14528                        if let Some(conv) = &w.conv {
14529                            meta[base + 10] = pack.put(conv);
14530                            meta[base + 24] = (conv.len() / self.hidden_size) as u32;
14531                        } else {
14532                            meta[base + 24] = 0;
14533                        }
14534                    }
14535                    AttnKind::Full { wq, wk, wv, wo, .. } => {
14536                        full_seen = true;
14537                        meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(wq));
14538                        meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(wk));
14539                        meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(wv));
14540                        meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(wo));
14541                        meta[base + 26] = (li * kv_stride) as u32;
14542                    }
14543                    AttnKind::Bounded(w) => {
14544                        meta[base + 11] = pack.put(&Self::embryo_qtensor_f32(&w.wq));
14545                        meta[base + 12] = pack.put(&Self::embryo_qtensor_f32(&w.wk));
14546                        meta[base + 13] = pack.put(&Self::embryo_qtensor_f32(&w.wv));
14547                        meta[base + 14] = pack.put(&Self::embryo_qtensor_f32(&w.wo));
14548                        // Ring slot of this anchor (anchors only, packed).
14549                        meta[base + 26] = (bounded_seen * kv_stride) as u32;
14550                        meta[base + 27] = pack.put(&w.sink_k);
14551                        meta[base + 28] = pack.put(&w.sink_v);
14552                        bounded_seen += 1;
14553                    }
14554                    _ => return None,
14555                }
14556                match &lw.ffn {
14557                    FfnKind::Dense(d) => {
14558                        meta[base + 15] = 0;
14559                        meta[base + 16] = 0;
14560                        meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&d.gate_proj));
14561                        meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&d.up_proj));
14562                        meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&d.down_proj));
14563                    }
14564                    FfnKind::Moe(m) => {
14565                        let r = m.resonance.as_ref().unwrap();
14566                        let (shared, _) = m.shared.as_ref().unwrap();
14567                        meta[base + 15] = 1;
14568                        meta[base + 16] = m.experts.len() as u32;
14569                        meta[base + 17] = pack.put(&r.mu);
14570                        meta[base + 18] = pack.put(&r.u);
14571                        meta[base + 19] = pack.put(&r.bias);
14572                        meta[base + 20] = r.k as u32;
14573                        // Word 30: the growth shell as the runtime applies
14574                        // it now (`+inf` on trunk rows, the stored finite
14575                        // shell on grown rows, all `+inf` under
14576                        // `CMF_GROWTH_SHELL=off`), followed by one `−∞`
14577                        // sentinel at index E the kernel writes as the
14578                        // score of an expert outside its shell — WGSL has
14579                        // no infinity literal, so the value travels as
14580                        // data (`embryo_core_route_finalize`).
14581                        let mut shell = r.effective_shell(m.experts.len());
14582                        shell.push(f32::NEG_INFINITY);
14583                        meta[base + 30] = pack.put(&shell);
14584                        meta[base + 21] = pack.put(&Self::embryo_qtensor_f32(&shared.gate_proj));
14585                        meta[base + 22] = pack.put(&Self::embryo_qtensor_f32(&shared.up_proj));
14586                        meta[base + 23] = pack.put(&Self::embryo_qtensor_f32(&shared.down_proj));
14587                        for (e, ex) in m.experts.iter().enumerate() {
14588                            meta[base + 32 + e * 3] =
14589                                pack.put(&Self::embryo_qtensor_f32(&ex.gate_proj));
14590                            meta[base + 33 + e * 3] =
14591                                pack.put(&Self::embryo_qtensor_f32(&ex.up_proj));
14592                            meta[base + 34 + e * 3] =
14593                                pack.put(&Self::embryo_qtensor_f32(&ex.down_proj));
14594                        }
14595                    }
14596                    FfnKind::DenseMoe(_) => return None,
14597                }
14598            }
14599            if !full_seen && self.num_layers == 0 {
14600                return None;
14601            }
14602            let id = {
14603                static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
14604                NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
14605            };
14606            let model = crate::gpu::EmbryoGraphModel {
14607                id,
14608                hidden: self.hidden_size,
14609                intermediate: self.intermediate_size,
14610                vocab: self.vocab_size,
14611                layers: self.num_layers,
14612                phase_heads: vmf.map(|c| c.num_heads).unwrap_or(0),
14613                nphase: vmf.map(|c| c.nphase).unwrap_or(0),
14614                phase_dv: vmf.map(|c| c.value_head_dim).unwrap_or(0),
14615                anchor_q_heads: self.num_heads,
14616                anchor_kv_heads: self.num_kv_heads,
14617                anchor_head_dim: self.head_dim,
14618                rotary_dim: self.rotary_dim,
14619                max_seq: self.kv_cache.max_seq_len,
14620                cluster_count,
14621                cluster_size: self.vocab_size / cluster_count,
14622                phase_state_len: vmf.map(|c| c.state_len()).unwrap_or(0),
14623                state_stride,
14624                kv_stride,
14625                norm_gemma: matches!(self.norm_style, NormStyle::Gemma),
14626                phase_mass: vmf.map(|c| c.phase_mass).unwrap_or(0.0),
14627                weights: pack.data,
14628                meta,
14629                lm_head: Self::embryo_qtensor_f32(&self.weights.lm_head),
14630                clusters: clusters.as_ref().clone(),
14631                final_norm: self.weights.final_norm.clone(),
14632                inv_freq: self.inv_freq.as_ref().clone(),
14633                bounded: bounded.is_some(),
14634                kv_layers: if bounded.is_some() {
14635                    bounded_seen
14636                } else {
14637                    self.num_layers
14638                },
14639                state_layers: phase_seen + gdn_seen,
14640                anchor_window,
14641                anchor_sink,
14642                phase_layers: phase_seen,
14643                gdn_layers: gdn_seen,
14644                gdn_heads: gdn.map(|g| g.num_v_heads).unwrap_or(0),
14645                gdn_k_heads: gdn.map(|g| g.num_k_heads).unwrap_or(0),
14646                gdn_dk: gdn.map(|g| g.key_head_dim).unwrap_or(0),
14647                gdn_dv: gdn.map(|g| g.value_head_dim).unwrap_or(0),
14648                gdn_kk: gdn.map(|g| g.conv_kernel).unwrap_or(0),
14649            };
14650            self.embryo_graph = Some(std::sync::Arc::new(model));
14651        }
14652        self.embryo_graph.clone()
14653    }
14654
14655    fn forward_layers_span(
14656        &mut self,
14657        hidden: &[f32],
14658        position: usize,
14659        task_mask: Option<&TaskMask>,
14660        from: usize,
14661        upto: Option<usize>,
14662    ) -> Vec<f32> {
14663        debug_assert!(
14664            from == 0
14665                || (self.dsv4.is_none()
14666                    && self.dsv41.is_none()
14667                    && self.qwen4_exp.is_none()
14668                    && self.g3n.is_none())
14669        );
14670        // Every plain forward — the whole-token Metal graph (`q1_graph_gpu`
14671        // wraps the GDN owners zero-copy and reallocates them on a size
14672        // change) and the CPU layer loop (reads/swaps `linear_state`) —
14673        // must see the previous speculative commit's asynchronous replay
14674        // complete. One mutex probe when nothing is pending.
14675        #[cfg(target_os = "macos")]
14676        if !crate::gpu_metal::wait_replay() {
14677            self.fail_metal_graph("the pending async replay failed before a plain forward");
14678            return vec![0.0; self.hidden_size];
14679        }
14680        if let Some(b) = &mut self.qwen4_exp {
14681            let _ = (task_mask, upto);
14682            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14683            let mut logits = Vec::new();
14684            crate::qwen4_exp::forward_token(
14685                &b.0,
14686                &b.1,
14687                &b.2,
14688                &mut b.3,
14689                token_id,
14690                position,
14691                &self.inv_freq,
14692                self.pool.as_deref(),
14693                &mut logits,
14694                true,
14695            );
14696            self.graph_logits = Some(logits);
14697            return vec![0.0; self.hidden_size];
14698        }
14699        // DeepSeek-V4 runs its own stack: the state is hc_mult copies, and
14700        // the forward returns LOGITS, not a hidden — the head is inside it
14701        // (the final fold sits between the last layer and the norm). The
14702        // token id rides in `hidden[0]`, written by embed_single, because
14703        // the hash layers route by id rather than by content.
14704        if let Some(b) = &mut self.dsv4 {
14705            let _ = (task_mask, upto);
14706            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14707            let (g, layers, cfg, st) = (&b.0, &b.1, b.2, &mut b.3);
14708            st.pos = position;
14709            let mut logits = Vec::new();
14710            crate::dsv4::forward_token(
14711                g,
14712                layers,
14713                &cfg,
14714                st,
14715                token_id,
14716                &self.inv_freq,
14717                self.pool.as_deref(),
14718                &mut logits,
14719            );
14720            self.graph_logits = Some(logits);
14721            self.dspark_probe(position, token_id);
14722            // The caller expects a hidden; the logits went out of band, as
14723            // with the fused lm_head path.
14724            return vec![0.0; self.hidden_size];
14725        }
14726        // DeepSeek-V4.1 owns its complete stack and emits logits out of band.
14727        if let Some(b) = &mut self.dsv41 {
14728            let _ = (task_mask, upto);
14729            let token_id = hidden.first().copied().unwrap_or(0.0) as u32;
14730            let mut logits = Vec::new();
14731            crate::dsv41::forward_token(
14732                &b.0,
14733                &b.1,
14734                &b.2,
14735                &mut b.3,
14736                token_id,
14737                position,
14738                self.pool.as_deref(),
14739                &mut logits,
14740            );
14741            self.graph_logits = Some(logits);
14742            return vec![0.0; self.hidden_size];
14743        }
14744        // Gemma-3n runs its own stack (4 AltUp replicas don't fit this
14745        // loop); `hidden` is the extended embedding from embed_single.
14746        if let Some(b) = &self.g3n {
14747            let _ = (task_mask, upto);
14748            return crate::g3n::g3n_forward(
14749                &b.0,
14750                &b.1,
14751                hidden,
14752                position,
14753                &mut self.kv_cache.layers,
14754                self.num_heads,
14755                self.num_kv_heads,
14756                self.head_dim,
14757                self.pool.as_deref(),
14758            );
14759        }
14760        // Cortiq Embryo owns a separate resident graph: phase recurrent
14761        // state, resonance routing, the GQA anchor KV and hierarchical head
14762        // all execute in one Vulkan submit. It is limited to a complete
14763        // unmasked stack; spans and task masks retain the exact host path.
14764        // `CMF_EMBRYO_DBG=1` names the gate that keeps a token off the
14765        // resident graph — every refusal below is otherwise silent.
14766        if from == 0
14767            && upto.is_none()
14768            && task_mask.is_none()
14769            && self.anchor_core.is_some()
14770            && std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1")
14771        {
14772            static ONCE: std::sync::Once = std::sync::Once::new();
14773            ONCE.call_once(|| {
14774                eprintln!(
14775                    "embryo-dbg: graph_on decode={} prefill={} resident_env={:?} enabled_here={} \
14776                     unsupported={} eligible={}",
14777                    crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode),
14778                    crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill),
14779                    std::env::var("CMF_EMBRYO_RESIDENT").ok(),
14780                    crate::gpu::enabled_here(),
14781                    self.graph_refused(),
14782                    self.embryo_resident_eligible(),
14783                );
14784            });
14785        }
14786        if from == 0
14787            && upto.is_none()
14788            && task_mask.is_none()
14789            // Embryo's recurrent/KV state has no host import path.  Do not
14790            // seed it for a prefill-only graph and then silently decode from
14791            // an empty CPU cache; both phases must select the resident owner.
14792            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Decode)
14793            && crate::gpu::wgpu_graph_on(crate::gpu::GraphPhase::Prefill)
14794            // This whole-token Embryo path remains explicitly opt-in.
14795            // `CMF_GPU_WGPU_GRAPH=1` still enables the mature generic graph,
14796            // but must not silently select this model-specific resident path.
14797            && matches!(
14798                std::env::var("CMF_EMBRYO_RESIDENT").as_deref(),
14799                Ok("1") | Ok("parallel")
14800            )
14801            && crate::gpu::enabled_here()
14802            && !self.graph_refused()
14803            // Sequence owner: past position zero the device continues only
14804            // a sequence it holds. A host-owned sequence (the graph refused
14805            // at its start, or its prefix was prefilled on the host) keeps
14806            // the host path to its end — never a device attempt at p > 0
14807            // over an empty device image.
14808            && (position == 0 || self.device_sequence_position().is_some())
14809            && self.embryo_resident_eligible()
14810            && let Some(model) = self.ensure_embryo_graph()
14811        {
14812            let mut lg = Vec::new();
14813            if crate::gpu::forward_embryo_graph(&model, self.graph_kv_id, hidden, position, &mut lg)
14814            {
14815                self.graph_logits = Some(lg);
14816                return vec![0.0; self.hidden_size];
14817            }
14818            if std::env::var("CMF_EMBRYO_DBG").as_deref() == Ok("1") {
14819                eprintln!("embryo-dbg: forward_embryo_graph refused at position {position}");
14820            }
14821            // The refusal is THIS pipeline's: falling through once is safe
14822            // at position zero (the host owns the sequence from here),
14823            // while a refusal at a later position of a device-owned
14824            // sequence would mix a host KV/state path with a partial
14825            // device sequence.
14826            self.mark_graph_refused();
14827            if position != 0 {
14828                // Fail this sequence alone and leave nothing stale behind:
14829                // no reuse key, no host or device state for the next
14830                // request on this slot to "extend".
14831                self.kv_cache.clear();
14832                self.clear_history();
14833                crate::gpu::graph_kv_reset(self.graph_kv_id);
14834                panic!(
14835                    "resident Embryo graph refused at position {position}; refusing an in-flight host fallback"
14836                );
14837            }
14838        }
14839        let mut h = hidden.to_vec();
14840        // MiMo-V2 expert placement: decided before the graph or the per-op
14841        // arena can claim the budget the expert bank needs.
14842        self.mimo_moe_prepare();
14843        let _mimo_q8 = self.mimo_moe.is_on()
14844            .then(crate::qtensor::enter_full_gpu_q8_scope);
14845        // Split borrows: copy scalars / clone handles so the per-layer
14846        // cfg does not hold `&self` while the KV cache is `&mut`.
14847        let (nh, _nkv, _hd, hs, _rd, eps) = (
14848            self.num_heads,
14849            self.num_kv_heads,
14850            self.head_dim,
14851            self.hidden_size,
14852            self.rotary_dim,
14853            self.rms_eps,
14854        );
14855        let pool = self.pool.clone();
14856        // Opt-in wgpu token-graph attention (discrete Vulkan/DX12): the whole
14857        // attention sub-block runs resident in one submit. Off by default.
14858        // Whole-token wgpu graph: eligibility + arbitration.
14859        //  - explicit CMF_GPU_WGPU_GRAPH forces it on/off;
14860        //  - discrete adapters (4090: decode 76 -> 137 tok/s) and GDN
14861        //    hybrids (recurrent state device-resident, no CPU twin to
14862        //    race) TRUST it;
14863        //  - integrated/mobile adapters RACE it against the normal path
14864        //    at generation granularity (gpu::graph_race_*) — tiled
14865        //    mobile GPUs can turn the ~300-dispatch graph into seconds
14866        //    per token, while a fast phone GPU keeps its win.
14867        let graph_env = std::env::var("CMF_GPU_WGPU_GRAPH").ok();
14868        let graph_on = match graph_env.as_deref() {
14869            Some("0") => false,
14870            Some("prefill") => false, // decode keeps the per-op path
14871            Some(_) => true,
14872            // Unset: same discrete-only default as every other graph
14873            // site. "Is the GPU on" used to stand in here — which made
14874            // the 0.2 tok/s whole-token graph race-eligible on mobile
14875            // adapters and cost 12-14× on first tokens (cmfmobile
14876            // TUNING.md); integrated GPUs keep the per-op probe path.
14877            None => crate::gpu::wgpu_graph_default(),
14878        };
14879        let graph_trusted =
14880            graph_env.is_some() || crate::gpu::wgpu_graph_default() || self.gdn_cfg.is_some();
14881        let race_eligible = graph_on
14882            && upto.is_none()
14883            && task_mask.is_none()
14884            && from == 0
14885            && !self.graph_refused();
14886        let mut tail_start = 0usize;
14887        if race_eligible && crate::gpu::graph_race_use_graph(graph_trusted) {
14888            let t_graph = std::time::Instant::now();
14889            let mut lg = Vec::new();
14890            let mut gl = 0usize;
14891            let built = self.try_token_graph_wgpu(hidden, position, &mut lg, &mut gl);
14892            let declined = built.is_none();
14893            let built = match built {
14894                Some(Ok(hh)) => Some(hh),
14895                Some(Err(())) => {
14896                    // O(1) state was admitted before the device failure; the
14897                    // CPU mirrors are stale by construction.  Clear the whole
14898                    // sequence and stop rather than walking that stale state.
14899                    self.clear_sequence_state();
14900                    self.graph_failed
14901                        .store(true, std::sync::atomic::Ordering::Relaxed);
14902                    self.cancel
14903                        .store(true, std::sync::atomic::Ordering::Relaxed);
14904                    tracing::error!("token graph failed after admission; sequence state cleared");
14905                    return vec![0.0; self.hidden_size];
14906                }
14907                None => None,
14908            };
14909            // Past the transient guards (o1 still collecting, a softcap)
14910            // a refusal is about the weights and will never change —
14911            // remember it instead of walking every layer again next
14912            // token.
14913            if declined && !self.o1_active() && self.attn_softcap == 0.0 {
14914                self.mark_graph_refused();
14915            }
14916            graph_note(built.is_some(), gl, self.num_layers);
14917            if let Some(hh) = built {
14918                let dur = t_graph.elapsed();
14919                if std::env::var("CMF_GRAPH_PROF").is_ok() {
14920                    eprintln!("graph-call: {:.2} ms total", dur.as_secs_f64() * 1000.0);
14921                }
14922                if gl > 0 && gl < self.num_layers {
14923                    // Device prefix: the graph ran layers 0..gl and handed
14924                    // back the boundary hidden — the loop below owns the
14925                    // tail. The prefix layers' KV/state advanced on the
14926                    // device; the tail's advances on the host below. One
14927                    // boundary crossing per token.
14928                    h = hh;
14929                    tail_start = gl;
14930                } else if graph_trusted || !crate::gpu::graph_race_first_token_hopeless(dur) {
14931                    if !graph_trusted {
14932                        crate::gpu::graph_race_record(true, dur);
14933                    }
14934                    if !lg.is_empty() {
14935                        // Graph produced logits (final-norm + lm_head folded in) —
14936                        // pad/cap to vocab and hand them to the sampler directly.
14937                        lg.resize(self.vocab_size, 0.0);
14938                        if let Some(c) = self.final_softcap {
14939                            for l in lg.iter_mut() {
14940                                *l = c * (*l / c).tanh();
14941                            }
14942                        }
14943                        self.graph_logits = Some(lg);
14944                    }
14945                    return hh;
14946                }
14947                // Hopeless first graph token: discard it and fall through
14948                // to the normal path. Safe exactly here — the prompt KV is
14949                // still CPU-owned (chunked prefill), so recomputing this
14950                // position is exact; the mirror's extra row is never read
14951                // (the race just settled on the normal path).
14952            }
14953        }
14954        // KIMI-LINEAR HAS NO SPLIT BUG. The 2.6× reported from the
14955        // model rotation (12.2 tok/s on one card against 4.6 on two)
14956        // was a single measurement of a model whose arm arbitration is
14957        // borderline, and it did not survive repetition. Three runs an
14958        // arm, same binary, back to back:
14959        //   probe on : 1 GPU 9.5 / 5.7 / 5.9   2 GPU 7.8 / 13.0 / 13.3
14960        //   pinned   : 1 GPU 5.6 / 5.3 / 5.2   2 GPU 3.5 / 4.2 / 3.4
14961        // With the arms pinned the split costs about 1.45×, which is
14962        // what a layer split costs. With the probe free, TWO CARDS RUN
14963        // FASTER — because for this model the CPU arm wins some op
14964        // classes and the probe finds that.
14965        //
14966        // Two things do stand, and both are measured. The token graph
14967        // builds NOTHING here (`covered 0 of 14 layers [0..14)`), so
14968        // every layer walks per-op on either arm — that is where the
14969        // headroom is, not in the split. And this model's benchmark is
14970        // unusable without `CMF_GPU_PROBE=0`: the arbitration alone
14971        // moves it by more than 2×.
14972        //
14973        // Span runs (network split): the graph covers exactly [from..=upto]
14974        // — one submit per SEGMENT per token. No race: its state is global
14975        // and calibrated on full stacks, so spans take the graph only where
14976        // it is trusted by default (discrete adapters / CMF_GPU_WGPU_GRAPH).
14977        let span = from > 0 || upto.is_some();
14978        if span && graph_on && task_mask.is_none() && graph_trusted {
14979            let upto_excl = upto.map_or(self.num_layers, |u| u + 1);
14980            let mut lg = Vec::new();
14981            let mut gl = 0usize;
14982            let span_res =
14983                self.try_token_graph_wgpu_span(hidden, position, &mut lg, from, upto_excl, &mut gl);
14984            let span_res = match span_res {
14985                Some(Ok(hh)) => Some(hh),
14986                Some(Err(())) => {
14987                    self.clear_sequence_state();
14988                    self.graph_failed
14989                        .store(true, std::sync::atomic::Ordering::Relaxed);
14990                    self.cancel
14991                        .store(true, std::sync::atomic::Ordering::Relaxed);
14992                    tracing::error!(
14993                        "span token graph failed after admission; sequence state cleared"
14994                    );
14995                    return vec![0.0; self.hidden_size];
14996                }
14997                None => None,
14998            };
14999            graph_note(span_res.is_some(), gl, upto_excl - from);
15000            if std::env::var("CMF_GPU_DEBUG").is_ok() {
15001                // How much of the span the graph actually covered. A
15002                // prefix of nothing means every layer walks per-op and
15003                // the split's extra cost is elsewhere.
15004                static SEEN: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
15005                if SEEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed) < 4 {
15006                    eprintln!(
15007                        "span graph: covered {gl} of {} layers [{from}..{upto_excl}) res={}",
15008                        upto_excl - from,
15009                        span_res.is_some()
15010                    );
15011                }
15012            }
15013            if let Some(hh) = span_res {
15014                if gl == upto_excl - from {
15015                    if !lg.is_empty() {
15016                        lg.resize(self.vocab_size, 0.0);
15017                        if let Some(c) = self.final_softcap {
15018                            for l in lg.iter_mut() {
15019                                *l = c * (*l / c).tanh();
15020                            }
15021                        }
15022                        self.graph_logits = Some(lg);
15023                    }
15024                    crate::gpu::set_layer(-1);
15025                    return hh;
15026                }
15027                // Partial device prefix of the span: CPU owns the tail.
15028                h = hh;
15029                tail_start = from + gl;
15030            }
15031        }
15032        // Layers the host is about to run whose device mirror moved ahead
15033        // of the host cache (a device prefix that shrank since the prompt,
15034        // a batched-prefill prefix longer than this token's): bring their
15035        // rows over first. One comparison per layer when nothing lags.
15036        let t_race_cpu = (race_eligible && !graph_trusted).then(std::time::Instant::now);
15037
15038        // A partial graph is an explicit GPU-prefix / CPU-tail split. Keep
15039        // the tail PURE host-side: letting its QTensor hooks re-enter the
15040        // residency arena streams every omitted layer through Vulkan and the
15041        // driver's freed-allocation cache can grow to the full model size
15042        // (25.4 GiB observed with a 14 GiB budget on Granite 30B Q8_2F).
15043        // With a MiMo expert bank the tail is not a whole-layer host
15044        // stream: its experts run from the bank (never the arena) and its
15045        // projections stay per-op on the device, which the bank's placement
15046        // left room for.
15047        let host_tail = tail_start > from;
15048        let _host_tail = (host_tail && !self.mimo_moe.is_on()).then(crate::gpu::enter_cpu_scope);
15049        let automatic_gpu_prefix = self.automatic_gpu_prefix();
15050
15051        let _prof_layers = crate::cpuprof::time(crate::cpuprof::Slot::Layers);
15052        #[cfg(target_os = "macos")]
15053        let mut gpu_skip_until = 0usize;
15054        for li in tail_start.max(from)..self.num_layers {
15055            let _capacity_tail = automatic_gpu_prefix
15056                .filter(|&prefix| li >= prefix && !self.mimo_moe.is_dynamic(li, host_tail))
15057                .map(|_| crate::gpu::enter_cpu_scope());
15058            crate::gpu::set_layer(li as i64); // layer-split GPU/CPU (CMF_GPU_LAYERS)
15059            if let Some(u) = upto {
15060                if li > u {
15061                    break;
15062                }
15063            }
15064            if let Some(mask) = task_mask {
15065                if !mask.layer_alive(li) {
15066                    continue; // dead layer: residual pass-through
15067                }
15068            }
15069            // Whole-block q1 token graph: a run of consecutive q1
15070            // layers — GDN and full attention — executes with one sync
15071            // per CPU attend instead of per op (macOS/Metal).
15072            #[cfg(target_os = "macos")]
15073            {
15074                if li < gpu_skip_until {
15075                    continue;
15076                }
15077                if task_mask.is_none() {
15078                    let end = self.q1_graph_gpu(li, upto, position, &mut h);
15079                    if self
15080                        .graph_failed
15081                        .load(std::sync::atomic::Ordering::Relaxed)
15082                    {
15083                        // The graph may have mutated device state before a
15084                        // command-buffer error. Never continue with a CPU
15085                        // tail or read a stale host mirror after admission.
15086                        return vec![0.0; self.hidden_size];
15087                    }
15088                    if end > li {
15089                        gpu_skip_until = end;
15090                        // Looped Transformer: the graph stopped at a loop
15091                        // boundary — apply final norm before the next iteration.
15092                        if self.is_loop_end(end - 1) && end < self.num_layers {
15093                            h = inference::rms_norm(
15094                                &h,
15095                                &self.weights.final_norm,
15096                                self.rms_eps,
15097                                self.norm_style,
15098                            );
15099                        }
15100                        continue;
15101                    }
15102                }
15103            }
15104
15105            if task_mask.is_none() {
15106                match self.mimo_graph_layer_rows(li, &mut h, &[position]) {
15107                    crate::gpu::BatchGraphOutcome::Completed => continue,
15108                    crate::gpu::BatchGraphOutcome::Failed => return vec![0.0; self.hidden_size],
15109                    crate::gpu::BatchGraphOutcome::Declined => {},
15110                }
15111            }
15112            #[cfg(feature = "gpu")]
15113            self.pull_lagging_host_kv(li, li + 1, position);
15114            let lw = &self.weights.layers[self.phys_layer(li)];
15115            if let Ok(tp) = std::env::var("CMF_TRACE_POS") {
15116                if tp.parse::<usize>().ok() == Some(position) {
15117                    let n: f32 = h.iter().map(|x| x * x).sum::<f32>().sqrt();
15118                    eprintln!(
15119                        "TRACE pos {position} layer {li}: |h| = {n:.6} h0 {:.6} h1 {:.6}",
15120                        h[0], h[1]
15121                    );
15122                }
15123            }
15124            // Norm into the pipeline scratch — the returning rms_norm
15125            // allocated twice per layer per token (roadmap §3 P0).
15126            let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
15127            inference::rms_norm_into(
15128                &h,
15129                &lw.input_norm,
15130                self.rms_eps,
15131                self.norm_style,
15132                &mut self.ws.n1,
15133            );
15134            drop(prof);
15135
15136            let attn_out = match &lw.attn {
15137                AttnKind::Mla(w) => {
15138                    let inv_freq_l = self.layer_inv_freq(li);
15139                    let rs = self.layer_rope_scale(li);
15140                    let eps = self.rms_eps;
15141                    let pool = self.pool.clone();
15142                    mla_attention(
15143                        w,
15144                        &self.ws.n1,
15145                        &mut self.kv_cache.layers[li],
15146                        position,
15147                        &inv_freq_l,
15148                        rs,
15149                        eps,
15150                        pool.as_deref(),
15151                    )
15152                }
15153                AttnKind::Linear(w) => {
15154                    let cfg = self.vmf_cfg.expect("linear layer without vmf_cfg");
15155                    vmf_phase_forward(
15156                        &self.ws.n1,
15157                        w,
15158                        &cfg,
15159                        &mut self.kv_cache.layers[li].linear_state,
15160                        self.pool.as_deref(),
15161                    )
15162                }
15163                AttnKind::Kda(w) => {
15164                    let cfg = self.kda_cfg.expect("kda layer without kda_cfg");
15165                    crate::linear_core::kda_forward(
15166                        &self.ws.n1,
15167                        w,
15168                        &cfg,
15169                        &mut self.kv_cache.layers[li].linear_state,
15170                        self.pool.as_deref(),
15171                    )
15172                }
15173                AttnKind::LinearGdn(w) => {
15174                    let cfg = self.gdn_cfg.expect("gdn layer without gdn_cfg");
15175                    gdn_forward(
15176                        &self.ws.n1,
15177                        w,
15178                        &cfg,
15179                        &mut self.kv_cache.layers[li].linear_state,
15180                        self.pool.as_deref(),
15181                    )
15182                }
15183                AttnKind::ShortConv(w) => {
15184                    let cfg = self
15185                        .short_conv_cfg
15186                        .expect("short-conv layer without short_conv_cfg");
15187                    short_conv_forward(
15188                        &self.ws.n1,
15189                        w,
15190                        &cfg,
15191                        &mut self.kv_cache.layers[li].linear_state,
15192                        self.pool.as_deref(),
15193                    )
15194                }
15195                AttnKind::Bounded(w) => {
15196                    // Natively bounded anchor: insert into the ring, attend
15197                    // over sinks ∪ window. No position, nothing appended.
15198                    let rope = self
15199                        .bounded_rope
15200                        .clone()
15201                        .expect("bounded layer without an installed rotation table");
15202                    let cfg = crate::bounded::BoundedAttnCfg {
15203                        num_heads: self.num_heads,
15204                        num_kv_heads: self.num_kv_heads,
15205                        head_dim: self.head_dim,
15206                        hidden_size: hs,
15207                        scale: self.attn_scale,
15208                        rope: &rope,
15209                        pool: pool.as_deref(),
15210                    };
15211                    crate::bounded::bounded_attention(
15212                        &self.ws.n1,
15213                        w,
15214                        &mut self.kv_cache.layers[li],
15215                        &cfg,
15216                    )
15217                }
15218                AttnKind::Full {
15219                    wq,
15220                    wk,
15221                    wv,
15222                    wo,
15223                    q_norm,
15224                    k_norm,
15225                    output_gate,
15226                    softplus_gate,
15227                    bias,
15228                } if self.kv_cache.layers[li].o1_sealed() => {
15229                    // O(1) override: decode on the sealed Nyström state
15230                    // instead of the growing KV cache.
15231                    let inv_freq_l = self.layer_inv_freq(li);
15232                    let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15233                    let cfg = QwenAttnCfg {
15234                        num_heads: self.layer_num_heads(li),
15235                        num_kv_heads: nkv_l,
15236                        head_dim: hd_l,
15237                        hidden_size: hs,
15238                        position,
15239                        inv_freq: &inv_freq_l,
15240                        rotary_dim: rd_l,
15241                        scale: self.attn_scale,
15242                        softcap: self.attn_softcap,
15243                        window: None,
15244                        v_norm: self.attn_v_norm,
15245                        qk_norm_after_rope: self.qk_norm_after_rope,
15246                        gate_sigmoid: self.proj_gate_sigmoid,
15247                        q_norm: q_norm.as_deref(),
15248                        k_norm: k_norm.as_deref(),
15249                        output_gate: *output_gate,
15250                        softplus_gate: softplus_gate
15251                            .as_ref()
15252                            .map(|(gate, per_head)| (gate, *per_head)),
15253                        rope_scale: self.layer_rope_scale(li),
15254                        bias: bias
15255                            .as_ref()
15256                            .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
15257                        rms_eps: eps,
15258                        norm_style: self.norm_style,
15259                        pool: pool.as_deref(),
15260                        v_head_dim: self.layer_v_dim(li),
15261                    };
15262                    attention::qwen_attention_nystrom(
15263                        &self.ws.n1,
15264                        wq,
15265                        wk,
15266                        wv,
15267                        wo,
15268                        &mut self.kv_cache.layers[li],
15269                        &cfg,
15270                    )
15271                }
15272                AttnKind::Full {
15273                    wq,
15274                    wk,
15275                    wv,
15276                    wo,
15277                    q_norm,
15278                    k_norm,
15279                    output_gate,
15280                    softplus_gate,
15281                    bias,
15282                } => 'attn: {
15283                    // wgpu token-graph attention (opt-in): whole sub-block in
15284                    // one submit, device K/V mirror. q1 only, no gate/bias/mask.
15285                    // Its kernel has no window, sink or narrow-V slot and
15286                    // one mirror geometry: such models stay on the CPU attend.
15287                    let dropin_reason =
15288                        graph_on.then(|| self.graph_attn_decline_reason()).flatten();
15289                    if let Some(reason) = dropin_reason {
15290                        self.note_graph_decline("wgpu attn dropin", reason);
15291                    }
15292                    if graph_on
15293                        && dropin_reason.is_none()
15294                        && !*output_gate
15295                        && softplus_gate.is_none()
15296                        && self.attention_heads_per_layer.is_none()
15297                        && bias.is_none()
15298                        && task_mask.is_none()
15299                    {
15300                        let inv_freq_l = self.layer_inv_freq(li);
15301                        let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15302                        let gemma = self.norm_style == cortiq_core::NormStyle::Gemma;
15303                        if let (Some((gm, qi)), Some((_, ki)), Some((_, vi)), Some((_, oi))) = (
15304                            wq.mapped_q1(),
15305                            wk.mapped_q1(),
15306                            wv.mapped_q1(),
15307                            wo.mapped_q1(),
15308                        ) {
15309                            let gm = gm.clone();
15310                            let mut out = vec![0f32; hs];
15311                            let cache = &self.kv_cache.layers[li];
15312                            if crate::gpu::attn_dropin(
15313                                &gm,
15314                                self.graph_kv_id,
15315                                li,
15316                                &self.ws.n1,
15317                                qi,
15318                                ki,
15319                                vi,
15320                                oi,
15321                                q_norm.as_deref(),
15322                                k_norm.as_deref(),
15323                                self.qk_norm_after_rope,
15324                                &inv_freq_l,
15325                                nh,
15326                                nkv_l,
15327                                hd_l,
15328                                rd_l,
15329                                hs,
15330                                position,
15331                                self.kv_cache.max_seq_len,
15332                                gemma,
15333                                eps as f32,
15334                                cache.k_heads(),
15335                                cache.v_heads(),
15336                                &mut out,
15337                            ) {
15338                                break 'attn out;
15339                            }
15340                        }
15341                    }
15342                    let masked = task_mask
15343                        .map(|m| m.head_flags(li, self.num_heads).iter().any(|&a| !a))
15344                        .unwrap_or(false);
15345                    let f32_view = (wq.as_f32(), wk.as_f32(), wv.as_f32(), wo.as_f32());
15346                    // The masked kernel knows one pipeline-wide geometry and
15347                    // RoPE table, no window and no sink.
15348                    let plain = self.layer_attn_plain(li);
15349                    match (masked, f32_view) {
15350                        // Historical masked path (f32 slices; the loader
15351                        // keeps masked models in f32).
15352                        (true, (Some(q), Some(k), Some(v), Some(o))) if plain => {
15353                            let active_heads = task_mask.unwrap().head_flags(li, self.num_heads);
15354                            attention::multi_head_attention(
15355                                &self.ws.n1,
15356                                q,
15357                                k,
15358                                v,
15359                                o,
15360                                &mut self.kv_cache.layers[li],
15361                                self.num_heads,
15362                                self.num_kv_heads,
15363                                self.head_dim,
15364                                self.hidden_size,
15365                                position,
15366                                &active_heads,
15367                                &self.inv_freq,
15368                            )
15369                        }
15370                        (masked, _) => {
15371                            if masked {
15372                                tracing::warn!(
15373                                    "layer {li}: head mask on quantized weights or on a \
15374                                     window/sink/per-layer-geometry layer not supported \
15375                                     yet — executing dense"
15376                                );
15377                            }
15378                            let inv_freq_l = self.layer_inv_freq(li);
15379                            let (nkv_l, hd_l, rd_l) = self.layer_geom(li);
15380                            let cfg = QwenAttnCfg {
15381                                num_heads: self.layer_num_heads(li),
15382                                num_kv_heads: nkv_l,
15383                                head_dim: hd_l,
15384                                hidden_size: hs,
15385                                position,
15386                                inv_freq: &inv_freq_l,
15387                                rotary_dim: rd_l,
15388                                scale: self.attn_scale,
15389                                softcap: self.attn_softcap,
15390                                window: self.layer_window(li),
15391                                v_norm: self.attn_v_norm,
15392                                qk_norm_after_rope: self.qk_norm_after_rope,
15393                                gate_sigmoid: self.proj_gate_sigmoid,
15394                                q_norm: q_norm.as_deref(),
15395                                k_norm: k_norm.as_deref(),
15396                                output_gate: *output_gate,
15397                                softplus_gate: softplus_gate
15398                                    .as_ref()
15399                                    .map(|(gate, per_head)| (gate, *per_head)),
15400                                rope_scale: self.layer_rope_scale(li),
15401                                bias: bias
15402                                    .as_ref()
15403                                    .map(|(a, b, c)| (a.as_slice(), b.as_slice(), c.as_slice())),
15404                                rms_eps: eps,
15405                                norm_style: self.norm_style,
15406                                pool: pool.as_deref(),
15407                                v_head_dim: self.layer_v_dim(li),
15408                            };
15409                            attention::qwen_attention(
15410                                &self.ws.n1,
15411                                wq,
15412                                wk,
15413                                wv,
15414                                wo,
15415                                &mut self.kv_cache.layers[li],
15416                                &cfg,
15417                            )
15418                        }
15419                    }
15420                }
15421            };
15422            // Gemma sandwich norm: normalize the attention branch before
15423            // it joins the residual stream.
15424            let attn_out = match &self.weights.layers[self.phys_layer(li)].attn_out_norm {
15425                Some(w) => inference::rms_norm(&attn_out, w, self.rms_eps, self.norm_style),
15426                None => attn_out,
15427            };
15428            let lw = &self.weights.layers[self.phys_layer(li)];
15429            let prof = crate::cpuprof::time(crate::cpuprof::Slot::Norms);
15430            inference::add_rmsnorm_fused_into(
15431                &mut h,
15432                &attn_out,
15433                &lw.post_norm,
15434                self.rms_eps,
15435                self.norm_style,
15436                &mut self.ws.p1,
15437            );
15438            drop(prof);
15439            let mut attn_out = attn_out;
15440            attention::recycle_buf(&mut attn_out);
15441            let post_normed = &self.ws.p1;
15442
15443            let ffn_masked = task_mask
15444                .map(|m| m.ffn_active_count(li) < self.intermediate_size)
15445                .unwrap_or(false);
15446            // One masked dense CONTRACT, dispatched by cost. The
15447            // activation-zeroing arm (the batched sweep's, validated
15448            // against the replica to 0.8%) computes the FULL fused FFN
15449            // and zeroes the dead — right whenever most neurons live.
15450            // The sparse arm reads ONLY active rows and down columns —
15451            // per-row dots are slower per element than the fused kernel,
15452            // so it pays only once the mask is deep enough. The 0.5
15453            // crossover is first-principles (fused kernels run ~2x the
15454            // per-row dot throughput); a shallow specialist (95% alive)
15455            // stays fused, a --target-sparsity bake flips arms on its
15456            // own weight.
15457            let ffn_out = match (ffn_masked, &lw.ffn) {
15458                // A defragged tube layer answers its own mask: the core
15459                // always runs, each tube runs when its bit is on, and
15460                // the tubes that are off are never read from the mmap.
15461                (_, FfnKind::Dense(d)) if !d.segs.is_empty() => {
15462                    let row = task_mask
15463                        .and_then(|tm| tm.ffn_masks.get(li))
15464                        .map(|v| v.as_slice());
15465                    tube_ffn(d, post_normed, 1, self.pool.as_deref(), row)
15466                }
15467                (true, FfnKind::Dense(d)) => {
15468                    let tm = task_mask.unwrap();
15469                    let alive = tm.ffn_active_count(li);
15470                    let deep = alive * 2 <= self.intermediate_size;
15471                    if deep && d.down_proj.sparse_col_ok() && !d.gate_proj.has_prism_contract() {
15472                        let active = tm.ffn_active_indices(li);
15473                        sparse_ffn_quant(
15474                            d,
15475                            post_normed,
15476                            &active,
15477                            self.hidden_size,
15478                            self.pool.as_deref(),
15479                        )
15480                    } else if deep
15481                        && let (Some(g), Some(u), Some(dn)) = (
15482                            d.gate_proj.as_f32(),
15483                            d.up_proj.as_f32(),
15484                            d.down_proj.as_f32(),
15485                        )
15486                    {
15487                        let active = tm.ffn_active_indices(li);
15488                        inference::sparse_ffn_forward(
15489                            post_normed,
15490                            g,
15491                            u,
15492                            dn,
15493                            self.hidden_size,
15494                            self.intermediate_size,
15495                            &active,
15496                            self.pool.as_deref(),
15497                        )
15498                    } else {
15499                        let row = tm.ffn_masks.get(li).map(|v| v.as_slice());
15500                        dense_ffn_batch(d, post_normed, 1, self.pool.as_deref(), row)
15501                    }
15502                }
15503                (true, FfnKind::Moe(m)) => {
15504                    // MoE is sparse by expert selection; a task mask
15505                    // narrows the ROUTABLE set via its expert fields
15506                    // (spec §5) when it carries them.
15507                    let allowed = task_mask.and_then(|tm| tm.expert_flags(li, m.experts.len()));
15508                    ffn_forward(
15509                        &lw.ffn,
15510                        post_normed,
15511                        self.pool.as_deref(),
15512                        allowed.as_deref(),
15513                    )
15514                }
15515                (true, FfnKind::DenseMoe(dm)) => dense_moe_ffn(
15516                    dm,
15517                    post_normed,
15518                    &h,
15519                    self.rms_eps,
15520                    self.norm_style,
15521                    self.pool.as_deref(),
15522                ),
15523                (false, _) => match &lw.ffn {
15524                    FfnKind::DenseMoe(dm) => dense_moe_ffn(
15525                        dm,
15526                        post_normed,
15527                        &h,
15528                        self.rms_eps,
15529                        self.norm_style,
15530                        self.pool.as_deref(),
15531                    ),
15532                    FfnKind::Moe(m)
15533                        if task_mask.is_none() && self.mimo_moe.is_dynamic(li, host_tail) =>
15534                    {
15535                        moe_ffn_banked(&mut self.mimo_moe, li, m, post_normed, self.pool.as_deref())
15536                    }
15537                    _ => {
15538                        let allowed = match (&lw.ffn, task_mask) {
15539                            (FfnKind::Moe(m), Some(tm)) => tm.expert_flags(li, m.experts.len()),
15540                            _ => None,
15541                        };
15542                        ffn_forward(
15543                            &lw.ffn,
15544                            post_normed,
15545                            self.pool.as_deref(),
15546                            allowed.as_deref(),
15547                        )
15548                    }
15549                },
15550            };
15551            let ffn_out = match &self.weights.layers[self.phys_layer(li)].ffn_out_norm {
15552                Some(w) => inference::rms_norm(&ffn_out, w, self.rms_eps, self.norm_style),
15553                None => ffn_out,
15554            };
15555            for (i, &f) in ffn_out.iter().enumerate() {
15556                h[i] += f;
15557            }
15558            let mut ffn_out = ffn_out;
15559            attention::recycle_buf(&mut ffn_out);
15560
15561            // Gemma-4: the layer output is scaled by a learned scalar.
15562            if let Some(sc) = self.weights.layers[self.phys_layer(li)].layer_scale {
15563                for v in h.iter_mut() {
15564                    *v *= sc;
15565                }
15566            }
15567            // CMF_LAYER_DUMP: this position's hidden after layer li.
15568            if self.layer_dump.is_some() {
15569                self.dump_layer_row(position, li, &h);
15570            }
15571
15572            // Looped Transformer: apply final norm at the end of each loop iteration.
15573            // Nanbeige 4.2: after layer 21 (virtual), apply norm before looping back to layer 0.
15574            if self.is_loop_end(li) && li + 1 < self.num_layers {
15575                h = inference::rms_norm(
15576                    &h,
15577                    &self.weights.final_norm,
15578                    self.rms_eps,
15579                    self.norm_style,
15580                );
15581            }
15582
15583            // Dynamic routing φ capture (on-policy): the
15584            // EMA of the post-residual hidden at the router's phi_layer,
15585            // updated as the context evolves during decode.
15586            if self.dyn_phi_layer == Some(li) {
15587                self.update_dyn_phi(&h);
15588            }
15589        }
15590        crate::gpu::set_layer(-1); // layers done — lm_head outside layer-split
15591        if let Some(t) = t_race_cpu {
15592            crate::gpu::graph_race_record(false, t.elapsed());
15593        }
15594
15595        h
15596    }
15597
15598    /// EMA of φ at the router layer (rolling, weight 0.2 = ~5-token
15599    /// horizon). First observation seeds it exactly.
15600    fn update_dyn_phi(&mut self, h: &[f32]) {
15601        const A: f32 = 0.2;
15602        if self.dyn_phi_ema.len() != h.len() {
15603            self.dyn_phi_ema = vec![0.0; h.len()];
15604            self.dyn_phi_seen = 0;
15605        }
15606        if self.dyn_phi_seen == 0 {
15607            self.dyn_phi_ema.copy_from_slice(h);
15608        } else {
15609            for (e, &v) in self.dyn_phi_ema.iter_mut().zip(h) {
15610                *e = (1.0 - A) * *e + A * v;
15611            }
15612        }
15613        self.dyn_phi_seen += 1;
15614    }
15615
15616    /// Current router φ (EMA at phi_layer); empty until first capture.
15617    pub fn dyn_phi(&self) -> &[f32] {
15618        &self.dyn_phi_ema
15619    }
15620
15621    /// Enable/disable φ capture at the router layer, reset the EMA.
15622    pub fn set_dyn_phi_layer(&mut self, layer: Option<usize>) {
15623        self.dyn_phi_layer = layer;
15624        self.dyn_phi_ema.clear();
15625        self.dyn_phi_seen = 0;
15626    }
15627
15628    /// Skills eligible for dynamic switching: (index, id, phi_layer).
15629    pub fn dynamic_skills(&self) -> Vec<(usize, String, usize)> {
15630        let Some(model) = &self.model else {
15631            return Vec::new();
15632        };
15633        model
15634            .header
15635            .skills
15636            .iter()
15637            .enumerate()
15638            .filter_map(|(i, sk)| {
15639                // A v2 record routes only through the request-level
15640                // backbone-gated decision (its status and gate are
15641                // checked there), never per token.
15642                if sk.is_v2() {
15643                    return None;
15644                }
15645                let ok = matches!(self.dyn_skill_layers.get(i), Some(Some(_)));
15646                let sel = sk.selection.as_ref()?;
15647                (ok).then(|| (i, sk.id.clone(), sel.phi_layer))
15648            })
15649            .collect()
15650    }
15651
15652    /// Index of the currently overlaid skill (None = backbone).
15653    pub fn active_skill(&self) -> Option<usize> {
15654        self.dyn_active
15655    }
15656
15657    /// Enable dynamic per-token skill routing: build the hysteresis
15658    /// router from the container's routable skills, start φ capture at
15659    /// their (shared) phi_layer. Returns the number of routable skills
15660    /// (0 = nothing to route; router stays off). Idempotent.
15661    pub fn enable_dynamic_routing(&mut self) -> usize {
15662        use crate::swarm::{DynRouter, RoutableSkill};
15663        let Some(model) = self.model.clone() else {
15664            return 0;
15665        };
15666        // Router policy v2 (spec §9.4) routes per REQUEST: the backbone is
15667        // the default and only the backbone-gated decision may pick a
15668        // skill. A per-token switch would bypass that gate (and change
15669        // the O(1) state mid-sequence), so the hysteresis router never
15670        // runs on such a file; the caller keeps the request-level
15671        // decision.
15672        if let Some(r) = &model.header.router {
15673            tracing::warn!(
15674                "dynamic routing disabled: this file declares router policy '{}' with \
15675                 granularity \"{}\" — the request-level decision applies instead",
15676                r.policy,
15677                r.granularity
15678            );
15679            return 0;
15680        }
15681        // Format-v2 skill records (bit SKILLS_V2) without a router policy:
15682        // their status/gate contract ("auto-routing requires active +
15683        // measured") lives in the backbone-gated decision only — the
15684        // hysteresis router would switch into a quarantined record
15685        // (fail-open). Refuse the whole file, not just its v2 records.
15686        if model.required_features & cortiq_core::format::features::SKILLS_V2 != 0
15687            || model.header.skills.iter().any(|s| s.is_v2())
15688        {
15689            tracing::warn!(
15690                "dynamic routing disabled: this file carries format-v2 skill records \
15691                 (SKILLS_V2) — they route per request through a router policy only"
15692            );
15693            return 0;
15694        }
15695        // A blend materialized f32 working tensors into the layers; there
15696        // is no single skill index to revert from → refuse (honest).
15697        if self.dyn_blend_loaded {
15698            tracing::warn!("dynamic routing unavailable on a blend-loaded pipeline");
15699            return 0;
15700        }
15701        // A statically-overlaid skill that is NOT FFN-eligible can't be
15702        // cheaply reverted at generation start → refuse rather than
15703        // silently keep it overlaid.
15704        if let Some(a) = self.dyn_active {
15705            if !matches!(self.dyn_skill_layers.get(a), Some(Some(_))) {
15706                tracing::warn!("loaded skill is not FFN-eligible — dynamic routing unavailable");
15707                return 0;
15708            }
15709        }
15710        let hidden = self.hidden_size;
15711        let mut skills = Vec::new();
15712        for (idx, id, _phi) in self.dynamic_skills() {
15713            if let Some(sel) = model.header.skills[idx].selection.as_ref() {
15714                if let Some(rs) = RoutableSkill::from_descriptor(idx, id, sel, hidden) {
15715                    skills.push(rs);
15716                }
15717            }
15718        }
15719        if skills.is_empty() {
15720            return 0;
15721        }
15722        // Skills should share a phi_layer; warn (not fail) if they don't.
15723        let phi = skills[0].phi_layer;
15724        if skills.iter().any(|s| s.phi_layer != phi) {
15725            tracing::warn!("routable skills disagree on phi_layer; using {phi}");
15726        }
15727        let n = skills.len();
15728        self.set_dyn_phi_layer(Some(phi));
15729        self.dyn_router = Some(DynRouter::new(skills));
15730        n
15731    }
15732
15733    /// Human-readable switch log from the last dynamic-routed generation.
15734    pub fn route_switches(&self) -> Vec<(usize, Option<String>, Option<String>)> {
15735        self.dyn_router
15736            .as_ref()
15737            .map(|r| r.switches.clone())
15738            .unwrap_or_default()
15739    }
15740
15741    /// LM head: hidden → logits [vocab_size]. The dominant matvec of
15742    /// every decode step — row-parallel on the worker pool.
15743    fn lm_head_forward(&self, hidden: &[f32]) -> Vec<f32> {
15744        let _mimo_q8 = self.mimo_moe.is_on()
15745            .then(crate::qtensor::enter_full_gpu_q8_scope);
15746        let rows = self.weights.lm_head.rows();
15747        let mut logits = attention::take_buf(rows.min(self.vocab_size));
15748        // Banked MiMo uses the same exact projection family for the
15749        // plain/draft head and the batched verification head. Read both
15750        // scale planes in-place instead of preparing per-op scale buffers.
15751        let served = self.mimo_moe.is_on() && crate::gpu::mimo_q8_short_enabled()
15752            && rows == self.vocab_size && !self.weights.lm_head.has_prism_contract()
15753            && self.weights.lm_head.graph_weight().is_some_and(|(model, idx, kind, _)| {
15754                kind == 7 && crate::gpu::q82_short_rows(model, idx, hidden, 1,
15755                    rows, self.hidden_size, &mut logits)
15756            });
15757        if !served {
15758            self.weights.lm_head.matvec(hidden, &mut logits, self.pool.as_deref());
15759        }
15760        logits.resize(self.vocab_size, 0.0);
15761        if let Some(m) = self.logit_multiplier {
15762            for l in logits.iter_mut() {
15763                *l *= m;
15764            }
15765        }
15766        if let Some(c) = self.final_softcap {
15767            for l in logits.iter_mut() {
15768                *l = c * (*l / c).tanh();
15769            }
15770        }
15771        if let Some(cm) = self.head_clusters.as_ref() {
15772            self.hierarchical_head_logprobs(hidden, cm, &mut logits);
15773        }
15774        logits
15775    }
15776
15777    /// Two-level head (Cortiq Embryo): in place, logits[v] ← log p(v) =
15778    /// (lc[c] − lse(lc)) + (logit[v] − lse over v's cluster block), c = v / S.
15779    fn hierarchical_head_logprobs(&self, hidden: &[f32], cm: &[f32], logits: &mut [f32]) {
15780        let h = hidden.len();
15781        let ncl = cm.len() / h.max(1);
15782        if ncl == 0 || logits.len() % ncl != 0 {
15783            return;
15784        }
15785        let cs = logits.len() / ncl;
15786        // cluster logits + log-softmax
15787        let mut lc = vec![0.0f32; ncl];
15788        for c in 0..ncl {
15789            let row = &cm[c * h..(c + 1) * h];
15790            let mut s = 0.0f32;
15791            for j in 0..h {
15792                s += row[j] * hidden[j];
15793            }
15794            lc[c] = s;
15795        }
15796        let mx = lc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15797        let lse: f32 = mx + lc.iter().map(|v| (v - mx).exp()).sum::<f32>().ln();
15798        for c in 0..ncl {
15799            let blk = &mut logits[c * cs..(c + 1) * cs];
15800            let bm = blk.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
15801            let bl: f32 = bm + blk.iter().map(|v| (v - bm).exp()).sum::<f32>().ln();
15802            let add = lc[c] - lse - bl;
15803            for v in blk.iter_mut() {
15804                *v += add;
15805            }
15806        }
15807    }
15808
15809    /// Prefill `ids` and return the next-token logits — what the model
15810    /// would predict next, WITHOUT committing to generation (introspection
15811    /// for `cortiq explain`). Clears and repopulates the KV cache; leaves
15812    /// the active overlay untouched.
15813    pub fn prefill_next_logits(&mut self, ids: &[u32], task_mask: Option<&TaskMask>) -> Vec<f32> {
15814        #[cfg(target_os = "macos")]
15815        crate::gpu_metal::set_io_namespace(self.graph_kv_id);
15816        self.clear_sequence_state();
15817        // This helper is used by the pooled classification endpoint, where
15818        // every request is a fresh sequence. The shared reset also clears the
15819        // wgpu token graph's device-side recurrent state.
15820        crate::gpu::graph_race_begin_generation();
15821        if task_mask.is_none() {
15822            self.o1_begin();
15823        }
15824        let mut hidden = vec![0.0f32; self.hidden_size];
15825        for (pos, &id) in ids.iter().enumerate() {
15826            let emb = self.embed_single(id);
15827            hidden = self.forward_layers(&emb, pos, task_mask);
15828        }
15829        if let Err(err) = self.o1_seal_checked() {
15830            self.o1_fail(err);
15831        }
15832        // Stacks that own their head (V4, V4.1, Qwen3.8-Flash-Next, GLM-5)
15833        // return a zero hidden and hand the logits out of band.
15834        if let Some(logits) = self.graph_logits.take() {
15835            return logits;
15836        }
15837        inference::rms_norm_into(
15838            &hidden,
15839            &self.weights.final_norm,
15840            self.rms_eps,
15841            self.norm_style,
15842            &mut self.ws.n1,
15843        );
15844        self.lm_head_forward(&self.ws.n1)
15845    }
15846}
15847
15848/// Convenience: deterministic tiny pipeline for tests.
15849pub fn create_test_pipeline(
15850    hidden_size: usize,
15851    intermediate_size: usize,
15852    num_heads: usize,
15853    num_kv_heads: usize,
15854    head_dim: usize,
15855    num_layers: usize,
15856    vocab_size: usize,
15857) -> Pipeline {
15858    // Small pseudo-random weights: constant weights make attention
15859    // degenerate and hide indexing bugs.
15860    let synth = |n: usize, salt: usize| -> Vec<f32> {
15861        (0..n)
15862            .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.2)
15863            .collect()
15864    };
15865    let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
15866        QTensor::from_f32(synth(rows * cols, salt), rows, cols)
15867    };
15868    let layer_weights: Vec<LayerWeights> = (0..num_layers)
15869        .map(|li| LayerWeights {
15870            input_norm: vec![1.0; hidden_size],
15871            post_norm: vec![1.0; hidden_size],
15872            attn_out_norm: None,
15873            ffn_out_norm: None,
15874            layer_scale: None,
15875            ffn: FfnKind::Dense(DenseFfn {
15876                gate_proj: qt(intermediate_size, hidden_size, li * 10 + 5),
15877                up_proj: qt(intermediate_size, hidden_size, li * 10 + 6),
15878                down_proj: qt(hidden_size, intermediate_size, li * 10 + 7),
15879                act: Act::Silu,
15880                down_t: None,
15881                segs: Vec::new(),
15882            }),
15883            attn: AttnKind::Full {
15884                bias: None,
15885                wq: qt(num_heads * head_dim, hidden_size, li * 10 + 1),
15886                wk: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 2),
15887                wv: qt(num_kv_heads * head_dim, hidden_size, li * 10 + 3),
15888                wo: qt(hidden_size, num_heads * head_dim, li * 10 + 4),
15889                q_norm: None,
15890                k_norm: None,
15891                output_gate: false,
15892                softplus_gate: None,
15893            },
15894        })
15895        .collect();
15896
15897    Pipeline::new(
15898        Tokenizer::byte_level(),
15899        PipelineWeights {
15900            embed_tokens: qt(vocab_size, hidden_size, 100),
15901            layers: layer_weights,
15902            lm_head: qt(vocab_size, hidden_size, 200),
15903            final_norm: vec![1.0; hidden_size],
15904        },
15905        hidden_size,
15906        intermediate_size,
15907        num_heads,
15908        num_kv_heads,
15909        head_dim,
15910        num_layers,
15911        num_layers, // physical_layers = num_layers (non-looped)
15912        false,      // loop_final_norm
15913        vocab_size,
15914        1e-6,
15915        10_000.0,
15916        NormStyle::Qwen,
15917        4096,
15918        SamplerConfig {
15919            seed: Some(42),
15920            ..Default::default()
15921        },
15922    )
15923}
15924
15925/// Batched dense-FFN: gate/up/down via matmat (element-wise the same
15926/// math as b × dense_ffn — the same dot kernels).
15927/// One mask bit, LSB-first per byte — `TaskMask::ffn_active_indices`'s
15928/// convention.
15929#[inline]
15930fn mask_bit(row: &[u8], j: usize) -> bool {
15931    (row.get(j >> 3).copied().unwrap_or(0) >> (j & 7)) & 1 != 0
15932}
15933
15934/// Zero the CLOSED neurons' activations in a [rows × inter] panel — the
15935/// masked-inference fast path's whole trick: full fused quant compute,
15936/// then the mask lands on the ACTIVATIONS, which is arithmetically the
15937/// pruned network without touching a quantized weight byte. Whole open
15938/// bytes (0xFF = 8 open neurons) skip in one test.
15939/// `CMF_FFN_MASK_GAIN` — Patent 12 FIG. 4, variance-preserving
15940/// rescaling: truncation removes a share of the layer's output energy,
15941/// so the survivors are scaled up to put the variance back where the
15942/// downstream norm expects it. A scalar here; per layer it is
15943/// `sqrt(total energy / kept energy)`.
15944fn mask_gain() -> f32 {
15945    static G: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
15946    *G.get_or_init(|| {
15947        std::env::var("CMF_FFN_MASK_GAIN")
15948            .ok()
15949            .and_then(|v| v.parse().ok())
15950            .unwrap_or(1.0)
15951    })
15952}
15953
15954fn zero_masked_cols(g: &mut [f32], rows: usize, inter: usize, row: &[u8]) {
15955    // With CMF_FFN_MEANFILL a closed neuron contributes its average
15956    // instead of nothing — same bytes read, one constant restored.
15957    let fill = meanfill().and_then(|(i, v)| {
15958        let li = crate::gpu::cur_layer();
15959        (*i == inter && li >= 0).then(|| &v[li as usize * inter..(li as usize + 1) * inter])
15960    });
15961    for r in 0..rows {
15962        let base = r * inter;
15963        for (bi, &byte) in row.iter().enumerate() {
15964            if byte == 0xFF {
15965                continue;
15966            }
15967            let j0 = bi * 8;
15968            for bit in 0..8 {
15969                let j = j0 + bit;
15970                if j < inter && byte & (1 << bit) == 0 {
15971                    g[base + j] = fill.map_or(0.0, |f| f[j]);
15972                }
15973            }
15974        }
15975    }
15976    let gain = mask_gain();
15977    if gain != 1.0 {
15978        for v in g[..rows * inter].iter_mut() {
15979            *v *= gain;
15980        }
15981    }
15982}
15983
15984/// True when neuron `i`'s bit is set (no mask = everything runs).
15985#[inline]
15986fn tube_bit(row: Option<&[u8]>, i: usize) -> bool {
15987    row.is_none_or(|r| mask_bit(r, i))
15988}
15989
15990/// Every bit below `n` set — the common case for a tube file's CORE,
15991/// where only the tube bits vary per task.
15992fn all_bits_on(row: &[u8], n: usize) -> bool {
15993    (0..n).all(|i| mask_bit(row, i))
15994}
15995
15996/// `CMF_TUBE_TOPK` — how many tubes a TOKEN may open (0 = the task mask
15997/// decides alone). This is the dense FFN read as a mixture: the tubes
15998/// are the experts a k-means over `gate_proj` rows found, and the token
15999/// picks among them. `CMF_TUBE_SCORE=gate` scores a tube by its own
16000/// gate (realizable: only `up`/`down` of the losers go unread),
16001/// `=oracle` scores by the true `silu(gate)·up` mass (the ceiling —
16002/// only `down` is saved, and the selection has read what it predicts).
16003fn tube_topk() -> usize {
16004    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
16005    *K.get_or_init(|| {
16006        std::env::var("CMF_TUBE_TOPK")
16007            .ok()
16008            .and_then(|v| v.parse().ok())
16009            .unwrap_or(0)
16010    })
16011}
16012
16013fn tube_score_oracle() -> bool {
16014    static O: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16015    *O.get_or_init(|| std::env::var("CMF_TUBE_SCORE").is_ok_and(|v| v == "oracle"))
16016}
16017
16018/// The routed arm of `tube_ffn`: a token opens only its best `k` tubes.
16019/// At `b == 1` (decode) the losers are genuinely never read — that is
16020/// the speed. At `b > 1` (the scoring sweep) every tube is computed and
16021/// the losers' activations are zeroed instead: same arithmetic, so the
16022/// perplexity is the routed model's, measured without a per-token
16023/// gather in the middle of a GEMM.
16024fn tube_ffn_routed(
16025    d: &DenseFfn,
16026    xs: &[f32],
16027    b: usize,
16028    pool: Option<&Pool>,
16029    mask_row: Option<&[u8]>,
16030    k: usize,
16031) -> Vec<f32> {
16032    let hidden = d.down_proj.rows();
16033    let core = d.gate_proj.rows();
16034    let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
16035    let mut out = match (b, core_full, mask_row) {
16036        (1, true, _) => dense_ffn(d, xs, pool),
16037        (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
16038        (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
16039        (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
16040    };
16041    let cand: Vec<usize> = (0..d.segs.len())
16042        .filter(|&i| tube_bit(mask_row, d.segs[i].start))
16043        .collect();
16044    if cand.is_empty() {
16045        return out;
16046    }
16047    // gate (and, where the score or the batch needs it, up) per tube.
16048    // The SCORE is taken at the point the serving path could take it:
16049    // off the gate alone, or off the finished activation for the oracle.
16050    let oracle = tube_score_oracle();
16051    let mut acts: Vec<Vec<f32>> = Vec::with_capacity(cand.len());
16052    let mut scores = vec![0f32; b * cand.len()];
16053    for (ci, &i) in cand.iter().enumerate() {
16054        let seg = &d.segs[i];
16055        let w = seg.width;
16056        let mut g = vec![0.0f32; b * w];
16057        if b == 1 {
16058            seg.gate.matvec(xs, &mut g, pool);
16059        } else {
16060            seg.gate.matmat(xs, b, &mut g, pool);
16061        }
16062        for v in g.iter_mut() {
16063            *v = Act::Silu.combine(*v, 1.0);
16064        }
16065        if !oracle {
16066            for t in 0..b {
16067                scores[t * cand.len() + ci] =
16068                    g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
16069            }
16070        }
16071        if oracle || b > 1 {
16072            let mut u = vec![0.0f32; b * w];
16073            if b == 1 {
16074                seg.up.matvec(xs, &mut u, pool);
16075            } else {
16076                seg.up.matmat(xs, b, &mut u, pool);
16077            }
16078            for (a, &v) in g.iter_mut().zip(u.iter()) {
16079                *a *= v;
16080            }
16081            if oracle {
16082                for t in 0..b {
16083                    scores[t * cand.len() + ci] =
16084                        g[t * w..(t + 1) * w].iter().map(|v| v * v).sum::<f32>();
16085                }
16086            }
16087        }
16088        acts.push(g);
16089    }
16090    // per-token scores and the winners
16091    let keep = k.min(cand.len());
16092    let mut scratch: Vec<f32> = Vec::new();
16093    for t in 0..b {
16094        let mut sc: Vec<(f32, usize)> = (0..cand.len())
16095            .map(|ci| (scores[t * cand.len() + ci], ci))
16096            .collect();
16097        sc.sort_unstable_by(|x, y| y.0.total_cmp(&x.0));
16098        let mut alive = vec![false; cand.len()];
16099        for &(_, ci) in sc.iter().take(keep) {
16100            alive[ci] = true;
16101        }
16102        if b > 1 {
16103            for (ci, a) in acts.iter_mut().enumerate() {
16104                if !alive[ci] {
16105                    let w = d.segs[cand[ci]].width;
16106                    a[t * w..(t + 1) * w].fill(0.0);
16107                }
16108            }
16109        } else {
16110            // decode: finish only the winners — the losers' up/down
16111            // (and, with the gate score, everything but their gate)
16112            // are never touched.
16113            for (ci, &i) in cand.iter().enumerate() {
16114                if !alive[ci] {
16115                    continue;
16116                }
16117                let seg = &d.segs[i];
16118                let w = seg.width;
16119                let g = &mut acts[ci];
16120                if !tube_score_oracle() {
16121                    scratch.clear();
16122                    scratch.resize(w, 0.0);
16123                    seg.up.matvec(xs, &mut scratch, pool);
16124                    for (a, &v) in g.iter_mut().zip(scratch.iter()) {
16125                        *a *= v;
16126                    }
16127                }
16128                let mut acc = vec![0.0f32; hidden];
16129                seg.down.matvec(g, &mut acc, pool);
16130                for (o, a) in out.iter_mut().zip(&acc) {
16131                    *o += *a;
16132                }
16133            }
16134        }
16135    }
16136    if b > 1 {
16137        for (ci, &i) in cand.iter().enumerate() {
16138            let seg = &d.segs[i];
16139            let mut acc = vec![0.0f32; b * hidden];
16140            seg.down.matmat(&acts[ci], b, &mut acc, pool);
16141            for (o, a) in out.iter_mut().zip(&acc) {
16142                *o += *a;
16143            }
16144        }
16145    }
16146    out
16147}
16148
16149/// FFN of a defragged tube layer: the always-on core plus the tubes the
16150/// task mask switches on. Each tube is a normal tensor triple, so the
16151/// same kernels run it and an inactive tube's bytes are never read —
16152/// that is the whole point of the defrag (a scattered mask cannot skip
16153/// bytes; a contiguous one is just a smaller matrix).
16154fn tube_ffn(
16155    d: &DenseFfn,
16156    xs: &[f32],
16157    b: usize,
16158    pool: Option<&Pool>,
16159    mask_row: Option<&[u8]>,
16160) -> Vec<f32> {
16161    if tube_topk() > 0 {
16162        return tube_ffn_routed(d, xs, b, pool, mask_row, tube_topk());
16163    }
16164    let hidden = d.down_proj.rows();
16165    let core = d.gate_proj.rows();
16166    let core_full = mask_row.is_none_or(|r| all_bits_on(r, core));
16167    let mut out = match (b, core_full, mask_row) {
16168        (1, true, _) => dense_ffn(d, xs, pool),
16169        (1, false, Some(row)) => dense_ffn_masked(d, xs, pool, row),
16170        (_, true, _) => dense_ffn_batch(d, xs, b, pool, None),
16171        (_, false, row) => dense_ffn_batch(d, xs, b, pool, row),
16172    };
16173    TUBE_SCRATCH.with(|sc| {
16174        let mut sc = sc.borrow_mut();
16175        let [g, u, acc] = &mut *sc;
16176        for seg in &d.segs {
16177            if !tube_bit(mask_row, seg.start) {
16178                continue;
16179            }
16180            let w = seg.width;
16181            g.resize(b * w, 0.0);
16182            if b == 1
16183                && d.act == Act::Silu
16184                && QTensor::matvec_silu_mul(&seg.gate, &seg.up, xs, g, pool)
16185            {
16186                // g holds silu(gate)·up.
16187            } else {
16188                u.resize(b * w, 0.0);
16189                if b == 1 {
16190                    QTensor::matvec_many([&seg.gate, &seg.up], xs, [g, u], pool);
16191                } else {
16192                    seg.gate.matmat(xs, b, g, pool);
16193                    seg.up.matmat(xs, b, u, pool);
16194                }
16195                for i in 0..b * w {
16196                    g[i] = d.act.combine(g[i], u[i]);
16197                }
16198            }
16199            acc.resize(b * hidden, 0.0);
16200            acc.fill(0.0);
16201            if b == 1 {
16202                seg.down.matvec(g, acc, pool);
16203            } else {
16204                seg.down.matmat(g, b, acc, pool);
16205            }
16206            for (o, a) in out.iter_mut().zip(acc.iter()) {
16207                *o += *a;
16208            }
16209        }
16210        out
16211    })
16212}
16213
16214thread_local! {
16215    /// gate / up / down-accumulator scratch for the tube loop — a tube
16216    /// runs once per layer per token, and a fresh Vec each time is a
16217    /// malloc per tube per layer per token.
16218    static TUBE_SCRATCH: std::cell::RefCell<[Vec<f32>; 3]> =
16219        const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new()]) };
16220}
16221
16222/// CMF_PREFILL_PROF: cumulative ns of the batched walk's attention and
16223/// FFN halves (all layers, all chunks).
16224static PREFILL_SPLIT: [std::sync::atomic::AtomicU64; 2] =
16225    [std::sync::atomic::AtomicU64::new(0), std::sync::atomic::AtomicU64::new(0)];
16226
16227/// The wgpu side of `CMF_PREFILL_PROF`: q|k|v GEMMs' enqueue and readback
16228/// wait, the chunk attend's readback wait (ms, cumulative) and the int8
16229/// GEMMs the `zi_mm` route took.
16230fn wgpu_prefill_counters() -> (f64, f64, f64, u64) {
16231    #[cfg(feature = "gpu")]
16232    {
16233        use std::sync::atomic::Ordering::Relaxed;
16234        let ms = |a: &std::sync::atomic::AtomicU64| a.load(Relaxed) as f64 / 1e6;
16235        (
16236            ms(&crate::gpu_wgpu::GEMM_MANY_NS[0]),
16237            ms(&crate::gpu_wgpu::GEMM_MANY_NS[1]),
16238            ms(&crate::gpu_wgpu::CHUNK_ATTEND_WAIT_NS),
16239            crate::gpu_wgpu::ZI_GEMM_CALLS.load(Relaxed),
16240        )
16241    }
16242    #[cfg(not(feature = "gpu"))]
16243    (0.0, 0.0, 0.0, 0)
16244}
16245
16246/// The wgpu prefill mirror attend's outcome counters
16247/// (`gpu_wgpu::MIRROR_EVENTS`).
16248fn wgpu_mirror_counters() -> [u64; 6] {
16249    #[cfg(feature = "gpu")]
16250    {
16251        crate::gpu_wgpu::MIRROR_EVENTS
16252            .each_ref()
16253            .map(|a| a.load(std::sync::atomic::Ordering::Relaxed))
16254    }
16255    #[cfg(not(feature = "gpu"))]
16256    [0; 6]
16257}
16258
16259fn prefill_prof_on() -> bool {
16260    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16261    *ON.get_or_init(|| std::env::var_os("CMF_PREFILL_PROF").is_some())
16262}
16263
16264fn dense_ffn_batch(
16265    d: &DenseFfn,
16266    xs: &[f32],
16267    b: usize,
16268    pool: Option<&Pool>,
16269    mask_row: Option<&[u8]>,
16270) -> Vec<f32> {
16271    let inter = d.gate_proj.rows();
16272    let hidden = d.down_proj.rows();
16273    // Fused on-device SwiGLU when the device is in play: three separate
16274    // `matmat` calls are three round trips per layer, and the gate/up
16275    // panels (b × inter — 22 MB each at a 512-token chunk) cross the bus
16276    // twice for nothing. The kernel already existed for the image DiT;
16277    // the LLM prefill was simply never wired to it. A task mask needs the
16278    // activations on the host between the halves, so it keeps the CPU
16279    // arm below.
16280    // SiLU and the exact GELU (Spark-X2.5) have device arms; `q4_ffn_act`
16281    // hands SiLU to the very entry points this used to call.
16282    let fused_act = d.act.graph_act();
16283    if mask_row.is_none()
16284        && fused_act.is_some()
16285        && b >= 32
16286        && crate::gpu::enabled_here()
16287        && !crate::gpu::mm_killed()
16288        // The refit pass needs this layer's activations on the host; the
16289        // fused chain keeps them on the device. Refusing it here costs
16290        // one round trip and keeps every GEMM on the card — the
16291        // alternative was running the whole calibration on the CPU.
16292        && refit_dir().is_none()
16293        // Same for the mass/hit probes. The accumulator at the bottom of
16294        // this function only sees `g` when `g` came back to the host, so
16295        // a fused batch would leave it summing nothing — a probe that
16296        // reports zeros rather than failing, which is worse.
16297        && !ffn_probe_active()
16298    {
16299        if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16300            d.gate_proj.mapped_q4t(),
16301            d.up_proj.mapped_q4t(),
16302            d.down_proj.mapped_q4t(),
16303        ) {
16304            let mut out = vec![0.0f32; b * hidden];
16305            let act = fused_act.expect("checked above");
16306            if crate::gpu::q4_ffn_act(model, w1, w3, w2, xs, b, hidden, inter, false, act, &mut out)
16307            {
16308                return out;
16309            }
16310        }
16311        // The q4tp twin (same kernel family, scale from the row ladder) —
16312        // the DiT has run it in production since the pipeline containers;
16313        // the LLM prefill was simply never wired to it, so a q4tp model's
16314        // prefill panels stayed on the CPU.
16315        if let (Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16316            d.gate_proj.mapped_q4tp(),
16317            d.up_proj.mapped_q4tp(),
16318            d.down_proj.mapped_q4tp(),
16319        ) {
16320            let mut out = vec![0.0f32; b * hidden];
16321            let act = fused_act.expect("checked above");
16322            if crate::gpu::q4_ffn_act(model, w1, w3, w2, xs, b, hidden, inter, true, act, &mut out)
16323            {
16324                return out;
16325            }
16326        }
16327        // Every other device codec — int8, or gate/up and down in different
16328        // codecs: the three GEMMs with the panels kept on the card. Without
16329        // it a q8_2f prefill read both b·inter panels home, folded them on
16330        // one host thread, and sent the result back for down.
16331        // CMF_FFN_KEEP=0 keeps the per-GEMM path (A/B).
16332        if let (true, Some((model, w1)), Some((_, w3)), Some((_, w2))) = (
16333            std::env::var("CMF_FFN_KEEP").as_deref() != Ok("0"),
16334            d.gate_proj.mapped_device_gemm(),
16335            d.up_proj.mapped_device_gemm(),
16336            d.down_proj.mapped_device_gemm(),
16337        ) {
16338            let mut out = vec![0.0f32; b * hidden];
16339            let act = fused_act.expect("checked above");
16340            if crate::gpu::ffn_act_keep(model, w1, w3, w2, xs, b, hidden, inter, act, &mut out) {
16341                return out;
16342            }
16343        }
16344    }
16345    let mut g = vec![0.0f32; b * inter];
16346    d.gate_proj.matmat(xs, b, &mut g, pool);
16347    let mut u = vec![0.0f32; b * inter];
16348    d.up_proj.matmat(xs, b, &mut u, pool);
16349    if gate_topk() > 0 && d.act == Act::Silu {
16350        for t in 0..b {
16351            let row = &mut g[t * inter..(t + 1) * inter];
16352            for v in row.iter_mut() {
16353                *v = Act::Silu.combine(*v, 1.0);
16354            }
16355            keep_top_k(row, gate_topk());
16356        }
16357        for i in 0..b * inter {
16358            g[i] *= u[i];
16359        }
16360    } else {
16361        for i in 0..b * inter {
16362            g[i] = d.act.combine(g[i], u[i]);
16363        }
16364    }
16365    if let Some(row) = mask_row {
16366        zero_masked_cols(&mut g, b, inter, row);
16367    }
16368    if oracle_topk() > 0 {
16369        for t in 0..b {
16370            keep_top_k(&mut g[t * inter..(t + 1) * inter], oracle_topk());
16371        }
16372    }
16373    let mut out = vec![0.0f32; b * hidden];
16374    d.down_proj.matmat(&g, b, &mut out, pool);
16375    if refit_dir().is_some() {
16376        let li = crate::gpu::cur_layer();
16377        if li >= 0 {
16378            refit_accumulate(li as usize, &g, b, inter, &out, hidden, pool);
16379        }
16380    }
16381    // The DTG-MA probe, on the batched path: one prefill sweep gives the
16382    // same per-neuron statistic the per-position probe does, and on a 27B
16383    // that is minutes instead of hours.
16384    FFN_PROBE.with(|pr| {
16385        if let Some(acc) = pr.borrow_mut().as_mut() {
16386            let li = crate::gpu::cur_layer();
16387            if li < 0 {
16388                return;
16389            }
16390            let Some(row) = acc.get_mut(li as usize) else {
16391                return;
16392            };
16393            let sq = probe_sq();
16394            for t in 0..b {
16395                for (a, &v) in row.iter_mut().zip(&g[t * inter..(t + 1) * inter]) {
16396                    *a += if sq {
16397                        (v as f64) * (v as f64)
16398                    } else {
16399                        (v as f64).abs()
16400                    };
16401                }
16402            }
16403        }
16404    });
16405    out
16406}
16407
16408/// Batched MoE-FFN: router batched, positions are GROUPED by expert —
16409/// an expert's weights are read once for all its positions in the chunk
16410/// (the main prefill-GEMM win on MoE: 960MB/token of 35B experts).
16411/// Accumulate per-channel activation energy for `CMF_RMS_TRACE`.
16412fn accumulate_act(m: &MoeFfn, xs: &[f32], b: usize) {
16413    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16414    static DUMP: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
16415    let on = *ON.get_or_init(|| std::env::var("CMF_RMS_TRACE").is_ok());
16416    let dump = *DUMP.get_or_init(|| std::env::var("CMF_ACT_DUMP").is_ok());
16417    if (!on && !dump) || b == 0 {
16418        return;
16419    }
16420    let hidden = xs.len() / b;
16421    if on {
16422        let mut acc = m.act_sq.borrow_mut();
16423        if acc.len() < hidden {
16424            acc.resize(hidden, 0.0);
16425        }
16426        for t in 0..b {
16427            let row = &xs[t * hidden..(t + 1) * hidden];
16428            for (a, &v) in acc.iter_mut().zip(row) {
16429                *a += (v as f64) * (v as f64);
16430            }
16431        }
16432    }
16433    if dump {
16434        // Cap the capture: the covariance needs a few thousand rows, and a
16435        // whole prefill of every layer would be gigabytes for no extra rank.
16436        let cap: usize = std::env::var("CMF_ACT_DUMP_ROWS")
16437            .ok()
16438            .and_then(|v| v.parse().ok())
16439            .unwrap_or(4096);
16440        let mut rows = m.act_rows.borrow_mut();
16441        if rows.len() < cap * hidden {
16442            let take = b.min((cap * hidden - rows.len()) / hidden.max(1));
16443            rows.extend_from_slice(&xs[..take * hidden]);
16444        }
16445    }
16446}
16447
16448/// Send-able cursor over a Vec-of-Vecs: each pool worker writes only its
16449/// own slots (disjoint by construction in the caller).
16450#[derive(Clone, Copy)]
16451struct SendVecs(*mut Vec<f32>);
16452unsafe impl Send for SendVecs {}
16453unsafe impl Sync for SendVecs {}
16454impl SendVecs {
16455    #[inline]
16456    fn at(self, i: usize) -> *mut Vec<f32> {
16457        unsafe { self.0.add(i) }
16458    }
16459}
16460
16461fn moe_ffn_batch(
16462    m: &MoeFfn,
16463    xs: &[f32],
16464    b: usize,
16465    hidden: usize,
16466    pool: Option<&Pool>,
16467    allowed: Option<&[bool]>,
16468) -> Vec<f32> {
16469    accumulate_act(m, xs, b);
16470    let ne = m.experts.len();
16471    let mut logits = vec![0.0f32; b * ne];
16472    match &m.resonance {
16473        Some(r) => {
16474            let hdim = xs.len() / b.max(1);
16475            for bi in 0..b {
16476                r.scores(
16477                    &xs[bi * hdim..(bi + 1) * hdim],
16478                    &mut logits[bi * ne..(bi + 1) * ne],
16479                );
16480            }
16481        }
16482        None => m.router.matmat(xs, b, &mut logits, pool),
16483    }
16484
16485    // Assignments: expert → [(position, weight)] — same routing as
16486    // moe_ffn, per position (see `moe_route`).
16487    let mut assign: Vec<Vec<(usize, f32)>> = vec![Vec::new(); ne];
16488    {
16489        let mut st = m.stats.borrow_mut();
16490        if st.len() < ne {
16491            st.resize(ne, 0);
16492        }
16493        for bi in 0..b {
16494            let (idx, p, wsum) = moe_route(&logits[bi * ne..(bi + 1) * ne], m, allowed);
16495            for &e in &idx {
16496                st[e] += 1;
16497                assign[e].push((bi, p[e] / wsum));
16498            }
16499        }
16500    }
16501
16502    let mut out = vec![0.0f32; b * hidden];
16503    let cols = m.experts[0].gate_proj.cols();
16504    let run_expert = |d: &DenseFfn, list: &[(usize, f32)], out: &mut [f32]| {
16505        let sb = list.len();
16506        let mut sub = vec![0.0f32; sb * cols];
16507        for (k, &(bi, _)) in list.iter().enumerate() {
16508            sub[k * cols..(k + 1) * cols].copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16509        }
16510        let eo = dense_ffn_batch(d, &sub, sb, pool, None);
16511        for (k, &(bi, w)) in list.iter().enumerate() {
16512            for i in 0..hidden {
16513                out[bi * hidden + i] += w * eo[k * hidden + i];
16514            }
16515        }
16516    };
16517    // Routed experts: the panels are TINY (b·top_k spread over every
16518    // expert — a few positions each), so a pool dispatch per expert is
16519    // pure barrier cost. Invert the parallelism: workers take WHOLE
16520    // experts (serial math inside), then one deterministic scatter in
16521    // expert order — the exact accumulation order the serial loop had.
16522    let active: Vec<usize> = (0..ne).filter(|&e| !assign[e].is_empty()).collect();
16523    if pool.is_some() && active.len() >= 8 {
16524        let mut panels: Vec<Vec<f32>> = vec![Vec::new(); active.len()];
16525        {
16526            let panel_ptr = SendVecs(panels.as_mut_ptr());
16527            // Capture only the expert table: `m` itself carries RefCell
16528            // stats and must not cross the pool boundary.
16529            let experts = &m.experts;
16530            let (active_r, assign_r) = (&active, &assign);
16531            let inherit_cpu = crate::gpu::inherit_cpu_scope();
16532            let run = |start: usize, end: usize| {
16533                let _cpu_scope = inherit_cpu();
16534                for ai in start..end {
16535                    let e = active_r[ai];
16536                    let list = &assign_r[e];
16537                    let sb = list.len();
16538                    let mut sub = vec![0.0f32; sb * cols];
16539                    for (k, &(bi, _)) in list.iter().enumerate() {
16540                        sub[k * cols..(k + 1) * cols]
16541                            .copy_from_slice(&xs[bi * cols..(bi + 1) * cols]);
16542                    }
16543                    // SAFETY: each worker owns a disjoint panels[ai].
16544                    unsafe {
16545                        *panel_ptr.at(ai) = dense_ffn_batch(&experts[e], &sub, sb, None, None);
16546                    }
16547                }
16548            };
16549            match pool {
16550                Some(p) => p.run_rows(active.len(), &run),
16551                None => run(0, active.len()),
16552            }
16553        }
16554        for (ai, &e) in active.iter().enumerate() {
16555            for (k, &(bi, w)) in assign[e].iter().enumerate() {
16556                let eo = &panels[ai][k * hidden..(k + 1) * hidden];
16557                for i in 0..hidden {
16558                    out[bi * hidden + i] += w * eo[i];
16559                }
16560            }
16561        }
16562    } else {
16563        for &e in &active {
16564            run_expert(&m.experts[e], &assign[e], &mut out);
16565        }
16566    }
16567    if let Some((se, gate)) = &m.shared {
16568        let all: Vec<(usize, f32)> = if let Some(gate) = gate {
16569            let mut gl = vec![0.0f32; b];
16570            gate.matmat(xs, b, &mut gl, pool);
16571            (0..b)
16572                .map(|bi| (bi, 1.0 / (1.0 + (-gl[bi]).exp())))
16573                .collect()
16574        } else {
16575            (0..b).map(|bi| (bi, 1.0)).collect()
16576        };
16577        run_expert(se, &all, &mut out);
16578    }
16579    out
16580}
16581
16582/// Decode-exact multi-token MoE — the MiMo speculative verify's FFN. Row
16583/// `r` of the result is bit-identical to `moe_ffn(m, x_r)` on the CPU
16584/// (`moe_ffn_cpu` → `moe_ffn_cpu_batched`): router matvec per row, the same
16585/// routing, the same int8 gate/up/SiLU and down terms
16586/// (`QTensor::moe_gate_up_rows` / `moe_down_rows`), and the row's experts
16587/// summed in ITS route order from 0. What the rows share is the weight
16588/// traffic: each routed expert is read once for every row that picked it.
16589/// (`moe_ffn_batch`, the prompt path, groups the same way but sums in
16590/// expert-index order and runs blocked kernels on wide groups — close, not
16591/// bit-equal to decode.) Any layer the kernels do not cover, or a device
16592/// that could answer `moe_ffn` itself, walks `moe_ffn` row by row.
16593fn moe_ffn_rows_exact(
16594    m: &MoeFfn,
16595    xs: &[f32],
16596    b: usize,
16597    hidden: usize,
16598    pool: Option<&Pool>,
16599) -> Vec<f32> {
16600    let mut out = vec![0.0f32; b * hidden];
16601    let per_row = |out: &mut [f32]| {
16602        for r in 0..b {
16603            let o = moe_ffn(m, &xs[r * hidden..(r + 1) * hidden], pool, None);
16604            out[r * hidden..(r + 1) * hidden].copy_from_slice(&o);
16605        }
16606    };
16607    let covered = !crate::gpu::enabled_here()
16608        && moe_batch_enabled()
16609        && m.shared.is_none()
16610        && m.resonance.is_none()
16611        && FFN_PROBE.with(|pr| pr.borrow().is_none())
16612        && m.experts.iter().all(|d| d.act == Act::Silu);
16613    if !covered {
16614        per_row(&mut out);
16615        return out;
16616    }
16617    let ne = m.experts.len();
16618    // Routing, row by row, exactly as `moe_ffn`.
16619    let mut routes: Vec<(Vec<usize>, Vec<f32>)> = Vec::with_capacity(b);
16620    for r in 0..b {
16621        let x = &xs[r * hidden..(r + 1) * hidden];
16622        accumulate_act(m, x, 1);
16623        let mut logits = vec![0.0f32; ne];
16624        m.router.matvec(x, &mut logits, pool);
16625        let (idx, p, wsum) = moe_route(&logits, m, None);
16626        {
16627            let mut st = m.stats.borrow_mut();
16628            if st.len() < ne {
16629                st.resize(ne, 0);
16630            }
16631            for &e in &idx {
16632                st[e] += 1;
16633            }
16634        }
16635        let w: Vec<f32> = idx
16636            .iter()
16637            .map(|&e| p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]))
16638            .collect();
16639        routes.push((idx, w));
16640    }
16641    if routes.iter().any(|(idx, _)| idx.is_empty()) {
16642        per_row(&mut out);
16643        return out;
16644    }
16645    // Group the (row, expert) picks by expert, in first-seen order.
16646    let mut experts: Vec<usize> = Vec::new();
16647    let mut groups: Vec<Vec<usize>> = Vec::new();
16648    for (r, (idx, _)) in routes.iter().enumerate() {
16649        for &e in idx {
16650            match experts.iter().position(|&x| x == e) {
16651                Some(g) => groups[g].push(r),
16652                None => {
16653                    experts.push(e);
16654                    groups.push(vec![r]);
16655                }
16656            }
16657        }
16658    }
16659    let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
16660    let inter = m.experts[experts[0]].gate_proj.rows();
16661    let pairs: Vec<(&QTensor, &QTensor)> = experts
16662        .iter()
16663        .map(|&e| (&m.experts[e].gate_proj, &m.experts[e].up_proj))
16664        .collect();
16665    let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
16666    if !QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut gs, pool) {
16667        per_row(&mut out);
16668        return out;
16669    }
16670    let downs: Vec<&QTensor> = experts.iter().map(|&e| &m.experts[e].down_proj).collect();
16671    let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
16672    let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; hidden]).collect();
16673    if !QTensor::moe_down_rows(&downs, &lens, &gs, &mut ds, pool) {
16674        per_row(&mut out);
16675        return out;
16676    }
16677    // Where each (row, expert) term landed in the flat pair list.
16678    let mut slot = std::collections::HashMap::with_capacity(n_pairs);
16679    let mut p = 0usize;
16680    for (g, &e) in experts.iter().enumerate() {
16681        for &r in &groups[g] {
16682            slot.insert((r, e), p);
16683            p += 1;
16684        }
16685    }
16686    for (r, (idx, w)) in routes.iter().enumerate() {
16687        let terms: Vec<(&[f32], f32)> = idx
16688            .iter()
16689            .zip(w)
16690            .map(|(&e, &we)| (ds[slot[&(r, e)]].as_slice(), we))
16691            .collect();
16692        let row = &mut out[r * hidden..(r + 1) * hidden];
16693        for (i, dst) in row.iter_mut().enumerate() {
16694            // `moe_down_many`'s per-row sum: from 0, in route order.
16695            let mut acc = 0f32;
16696            for (d, we) in &terms {
16697                acc += we * d[i];
16698            }
16699            *dst = acc;
16700        }
16701    }
16702    out
16703}
16704
16705thread_local! {
16706    /// gate/up activation scratch for the dense FFN paths (single uses
16707    /// two slots, the fused pair all four) — these were fresh
16708    /// intermediate-size Vecs on every layer of every token.
16709    static FFN_SCRATCH: std::cell::RefCell<[Vec<f32>; 4]> =
16710        const { std::cell::RefCell::new([Vec::new(), Vec::new(), Vec::new(), Vec::new()]) };
16711}
16712
16713/// Dense SwiGLU FFN through QTensor matvecs (any storage).
16714fn dense_ffn(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16715    // Per-token sparsity, when the file was built for it: gate first,
16716    // then only the chosen neurons' up/down rows leave the mmap.
16717    if gate_topk() > 0
16718        && let Some(out) = dense_ffn_dynamic(d, x, pool, gate_topk())
16719    {
16720        return out;
16721    }
16722    // Whole-FFN GPU submit (этап 4.2 increment): gate → silu·up → down
16723    // chained in ONE command buffer with the intermediate activations
16724    // resident on the device — 3 per-op polls become 1 per layer. The
16725    // moe_block backend already implements exactly this chain; a dense
16726    // FFN is one expert with weight 1. Runtime probe: the chain still
16727    // pays one submit+poll per layer — alternate it against the pure-CPU
16728    // FFN and keep whichever is faster on this machine.
16729    // q1 FFNs offload at any practical size: the q1 CPU kernel is
16730    // compute-bound, so the UMA threshold logic does not apply — the
16731    // probe measures and decides either way.
16732    // The fused GPU block has no descriptor-aware Prism path: it would either
16733    // consume an unrotated activation or decline after inspecting the mixed
16734    // q2tp/q4tp tensors.  Do not let that structural refusal enter the FFN
16735    // probe's CPU_ONLY scope; the ordinary body below dispatches each matrix
16736    // through QTensor::matvec, which owns the signed FWHT + affine q2tp route.
16737    let prism_body = d.gate_proj.has_prism_contract()
16738        || d.up_proj.has_prism_contract()
16739        || d.down_proj.has_prism_contract();
16740    if !prism_body
16741        && crate::gpu::enabled_here()
16742        && (d.gate_proj.rows() >= crate::gpu::min_rows() || d.gate_proj.is_q1())
16743    {
16744        let arm = if d.gate_proj.is_q1() && crate::gpu::q1_force() {
16745            crate::gpu::ProbeArm::Gpu
16746        } else {
16747            crate::gpu::probe_arm(crate::gpu::OpClass::Ffn)
16748        };
16749        match arm {
16750            crate::gpu::ProbeArm::Gpu => {
16751                let t0 = std::time::Instant::now();
16752                if let Some(out) = dense_ffn_gpu(d, x, pool) {
16753                    crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
16754                    return out;
16755                }
16756                // Declined: no timing exists, so say so. Silence here is
16757                // what left `ffn` undecided for 9000 calls and cost a
16758                // failed device attempt on half of them.
16759                crate::gpu::probe_note_decline(crate::gpu::OpClass::Ffn);
16760            }
16761            crate::gpu::ProbeArm::CpuTimed => {
16762                let t0 = std::time::Instant::now();
16763                let out = crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16764                crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
16765                return out;
16766            }
16767            crate::gpu::ProbeArm::Cpu => {
16768                return crate::gpu::cpu_scope(|| dense_ffn_cpu(d, x, pool));
16769            }
16770        }
16771    }
16772    dense_ffn_cpu(d, x, pool)
16773}
16774
16775/// The pure-CPU dense-FFN body (also the fallback of every GPU refusal).
16776fn dense_ffn_cpu(d: &DenseFfn, x: &[f32], pool: Option<&Pool>) -> Vec<f32> {
16777    let inter = d.gate_proj.rows();
16778    FFN_SCRATCH.with(|s| {
16779        let mut s = s.borrow_mut();
16780        let [g, u, ..] = &mut *s;
16781        g.resize(inter, 0.0);
16782        // Fused gate+up+silu: one dispatch, no separate silu pass.
16783        // Falls back to matvec_many + silu loop for unsupported dtypes.
16784        if gate_topk() > 0 {
16785            // Gate first, select, and only then pay for `up`: the
16786            // measurement arm computes both and zeroes the losers, which
16787            // is the same arithmetic.
16788            u.resize(inter, 0.0);
16789            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16790            for i in 0..inter {
16791                g[i] = Act::Silu.combine(g[i], 1.0);
16792            }
16793            keep_top_k(g, gate_topk());
16794            for i in 0..inter {
16795                g[i] *= u[i];
16796            }
16797        } else if d.act == Act::Silu && {
16798            let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16799            QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool)
16800        } {
16801            // g now holds silu(gate)·up directly.
16802        } else {
16803            u.resize(inter, 0.0);
16804            // Multi-matrix job: gate+up under one pool dispatch.
16805            let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnGateUp);
16806            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
16807            for i in 0..inter {
16808                g[i] = d.act.combine(g[i], u[i]);
16809            }
16810        }
16811        // DTG-MA bake probe (Patent 2): accumulate this layer's
16812        // per-neuron activation mass while a probe pass is active.
16813        // `CMF_FFN_PROBE_TOPK=k` switches the statistic from mass to a
16814        // HIT COUNT — how many tokens rank the neuron in their own top
16815        // k. Mass asks "how loud is this neuron overall", the count
16816        // asks "how often does this task actually need it", and the two
16817        // rank neurons differently whenever a few tokens are loud.
16818        FFN_PROBE.with(|pr| {
16819            if let Some(acc) = pr.borrow_mut().as_mut() {
16820                let li = crate::gpu::cur_layer();
16821                if li >= 0 {
16822                    if let Some(row) = acc.get_mut(li as usize) {
16823                        match probe_topk() {
16824                            0 if probe_sq() => {
16825                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16826                                    *a += (v as f64) * (v as f64);
16827                                }
16828                            }
16829                            0 if probe_signed() => {
16830                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16831                                    *a += v as f64;
16832                                }
16833                            }
16834                            0 => {
16835                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16836                                    *a += (v as f64).abs();
16837                                }
16838                            }
16839                            k => {
16840                                let n = g.len();
16841                                let k = k.min(n);
16842                                let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
16843                                let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
16844                                    b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
16845                                });
16846                                let thr = *kth;
16847                                for (a, &v) in row.iter_mut().zip(g.iter()) {
16848                                    if v.abs() >= thr {
16849                                        *a += 1.0;
16850                                    }
16851                                }
16852                            }
16853                        }
16854                    }
16855                }
16856            }
16857        });
16858        if oracle_topk() > 0 {
16859            keep_top_k(g, oracle_topk());
16860        }
16861        {
16862            let li = crate::gpu::cur_layer();
16863            if li >= 0 {
16864                adump_row(li as usize, g);
16865            }
16866        }
16867        let mut out = attention::take_buf(d.down_proj.rows());
16868        let _prof = crate::cpuprof::time(crate::cpuprof::Slot::FfnDown);
16869        d.down_proj.matvec(g, &mut out, pool);
16870        out
16871    })
16872}
16873
16874/// Online accumulators for the AWNP refit of a narrowed FFN.
16875///
16876/// The refit needs `Gss = A_SᵀA_S` and `YA = YᵀA_S` per layer, where `A_S`
16877/// are the calibration activations of the KEPT neurons and `Y` the full
16878/// FFN output. Both are small enough to hold; the thing that is not is
16879/// the activations they are built from — a 27B layer would dump a
16880/// gigabyte per thousand tokens. So they are accumulated as the
16881/// calibration runs and written once at the end.
16882///
16883/// `CMF_FFN_REFIT=<dir>` holds `support.<L>.u32` (a u32 count then the
16884/// kept indices) for every layer to accumulate; `CMF_FFN_REFIT_FROM/TO`
16885/// bound the layer span so the accumulators fit in RAM.
16886pub struct RefitAcc {
16887    pub support: Vec<u32>,
16888    pub gss: Vec<f32>,
16889    pub ya: Vec<f32>,
16890    pub hidden: usize,
16891    pub tokens: u64,
16892    /// Activations staged transposed ([ns, t] and [hidden, t]) until the
16893    /// batch is worth a GEMM. The product costs `ns²` to move and add
16894    /// REGARDLESS of how many tokens went into it, so folding 16 chunks
16895    /// into one call cuts that cost 16× — it was 15 TB of traffic per
16896    /// calibration pass at one call per 256 tokens.
16897    pub buf_g: Vec<f32>,
16898    pub buf_o: Vec<f32>,
16899    pub buf_t: usize,
16900}
16901
16902/// The product buffer is SHARED across layers — one 473 MB allocation,
16903/// not one per layer (that was 30 GB of nothing on a 64-layer model).
16904/// It lives under the same lock as the accumulators.
16905type RefitState = (std::collections::HashMap<usize, RefitAcc>, Vec<f32>);
16906
16907static REFIT: std::sync::OnceLock<Option<(String, std::sync::Mutex<RefitState>)>> =
16908    std::sync::OnceLock::new();
16909
16910/// Is an FFN probe accumulator installed on this thread? The fused GPU
16911/// FFN must decline while one is, or the probe silently measures zero.
16912fn ffn_probe_active() -> bool {
16913    FFN_PROBE.with(|p| p.borrow().is_some())
16914}
16915
16916fn refit_dir() -> Option<&'static (String, std::sync::Mutex<RefitState>)> {
16917    REFIT
16918        .get_or_init(|| {
16919            std::env::var("CMF_FFN_REFIT").ok().map(|d| {
16920                (
16921                    d,
16922                    std::sync::Mutex::new((std::collections::HashMap::new(), Vec::new())),
16923                )
16924            })
16925        })
16926        .as_ref()
16927}
16928
16929/// Accumulate one prefill panel into the layer's refit statistics.
16930fn refit_accumulate(
16931    li: usize,
16932    g: &[f32],
16933    b: usize,
16934    inter: usize,
16935    out: &[f32],
16936    hidden: usize,
16937    pool: Option<&Pool>,
16938) {
16939    let Some((dir, map)) = refit_dir() else {
16940        return;
16941    };
16942    static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
16943    let (from, to) = *SPAN.get_or_init(|| {
16944        let g = |k: &str, d: usize| {
16945            std::env::var(k)
16946                .ok()
16947                .and_then(|v| v.parse().ok())
16948                .unwrap_or(d)
16949        };
16950        (
16951            g("CMF_FFN_REFIT_FROM", 0),
16952            g("CMF_FFN_REFIT_TO", usize::MAX),
16953        )
16954    });
16955    if li < from || li > to {
16956        return;
16957    }
16958    let mut guard = map.lock().unwrap();
16959    let (map, shared) = &mut *guard;
16960    let acc = match map.entry(li) {
16961        std::collections::hash_map::Entry::Occupied(e) => e.into_mut(),
16962        std::collections::hash_map::Entry::Vacant(e) => {
16963            let path = format!("{dir}/support.{li}.u32");
16964            let Ok(bytes) = std::fs::read(&path) else {
16965                eprintln!("refit: no {path} — layer {li} skipped");
16966                return;
16967            };
16968            let n = u32::from_le_bytes(bytes[0..4].try_into().unwrap()) as usize;
16969            let support: Vec<u32> = bytes[4..4 + n * 4]
16970                .chunks_exact(4)
16971                .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
16972                .collect();
16973            eprintln!(
16974                "refit: layer {li} support {n} ({:.0} MB of accumulator)",
16975                (n * n + hidden * n) as f64 * 4.0 / 1e6
16976            );
16977            e.insert(RefitAcc {
16978                gss: vec![0.0; n * n],
16979                ya: vec![0.0; hidden * n],
16980                buf_g: Vec::new(),
16981                buf_o: Vec::new(),
16982                buf_t: 0,
16983                support,
16984                hidden,
16985                tokens: 0,
16986            })
16987        }
16988    };
16989    let ns = acc.support.len();
16990    // Stage this chunk transposed; the GEMM fires once the batch is full.
16991    let cap = refit_batch();
16992    if acc.buf_g.is_empty() {
16993        acc.buf_g = vec![0.0; ns * cap];
16994        acc.buf_o = vec![0.0; hidden * cap];
16995    }
16996    let take = b.min(cap - acc.buf_t);
16997    for t in 0..take {
16998        let col = acc.buf_t + t;
16999        for (j, &n) in acc.support.iter().enumerate() {
17000            acc.buf_g[j * cap + col] = g[t * inter + n as usize];
17001        }
17002        for h in 0..hidden {
17003            acc.buf_o[h * cap + col] = out[t * hidden + h];
17004        }
17005    }
17006    acc.buf_t += take;
17007    acc.tokens += take as u64;
17008    if acc.buf_t < cap {
17009        return;
17010    }
17011    let bt = acc.buf_t;
17012    acc.buf_t = 0;
17013    // The GEMM WRITES its C (it zeroes the accumulators it uses), so the
17014    // chunk product lands in scratch and is added on — the one thing that
17015    // silently turns a Gram over 13 000 tokens into a Gram over 256.
17016    // Both products are `C[n, m] += X[n, b] · Yᵀ[b, m]` with X and Y
17017    // stored row-major [·, b] — exactly `gemm_nt_f32`'s shape, so the
17018    // card does them when it is up (this is the whole calibration's
17019    // cost: O(|S|²) per token, 2.9 PFLOP for a 27B pass). The tiled CPU
17020    // loop stays as the fallback. Neither accumulates, so the product
17021    // lands in scratch and is added on.
17022    let RefitAcc {
17023        gss,
17024        ya,
17025        buf_g,
17026        buf_o,
17027        ..
17028    } = acc;
17029    let need = (ns * ns).max(hidden * ns);
17030    if shared.len() < need {
17031        shared.resize(need, 0.0);
17032    }
17033    let scratch = &mut shared[..];
17034    let _ = bt;
17035    if crate::gpu::gemm_nt_f32_transient(buf_g, buf_g, &mut scratch[..ns * ns], ns, cap, ns) {
17036        add_into(gss, &scratch[..ns * ns], pool);
17037        if crate::gpu::gemm_nt_f32_transient(
17038            buf_o,
17039            buf_g,
17040            &mut scratch[..hidden * ns],
17041            hidden,
17042            cap,
17043            ns,
17044        ) {
17045            add_into(ya, &scratch[..hidden * ns], pool);
17046        } else {
17047            accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
17048        }
17049    } else {
17050        accum_outer_t(gss, ns, ns, cap, buf_g, buf_g, pool);
17051        accum_outer_t(ya, hidden, ns, cap, buf_o, buf_g, pool);
17052    }
17053    // No zeroing: the batch is always filled exactly (cap is a multiple
17054    // of the prefill chunk), and a memset of 178 MB a layer would cost
17055    // more than the GEMM.
17056}
17057
17058/// `CMF_FFN_REFIT_BATCH` — tokens staged before each GEMM (default 4096).
17059fn refit_batch() -> usize {
17060    static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17061    *B.get_or_init(|| {
17062        std::env::var("CMF_FFN_REFIT_BATCH")
17063            .ok()
17064            .and_then(|v| v.parse().ok())
17065            .unwrap_or(4096)
17066    })
17067}
17068
17069/// `c[m, n] += Σ_t left[m, t]·right[n, t]` — both operands transposed,
17070/// the CPU fallback for the staged batch.
17071fn accum_outer_t(
17072    c: &mut [f32],
17073    m: usize,
17074    n: usize,
17075    b: usize,
17076    left: &[f32],
17077    right: &[f32],
17078    pool: Option<&Pool>,
17079) {
17080    let ptr = SendMut(c.as_mut_ptr());
17081    let body = |i: usize| {
17082        let ptr = &ptr;
17083        let row = unsafe { std::slice::from_raw_parts_mut(ptr.0.add(i * n), n) };
17084        for t in 0..b {
17085            let a = left[i * b + t];
17086            if a == 0.0 {
17087                continue;
17088            }
17089            for (j, o) in row.iter_mut().enumerate() {
17090                *o += a * right[j * b + t];
17091            }
17092        }
17093    };
17094    match pool {
17095        Some(p) if m > 1 => p.run_rows(m, &|s, e| {
17096            for i in s..e {
17097                body(i);
17098            }
17099        }),
17100        _ => {
17101            for i in 0..m {
17102                body(i);
17103            }
17104        }
17105    }
17106}
17107
17108/// `dst += src`, spread over the pool — at 118 M floats a layer this is
17109/// not a loop to leave on one core.
17110fn add_into(dst: &mut [f32], src: &[f32], pool: Option<&Pool>) {
17111    let n = dst.len().min(src.len());
17112    match pool {
17113        Some(p) if n >= 1 << 16 => {
17114            let ptr = SendMut(dst.as_mut_ptr());
17115            let f = |s: usize, e: usize| {
17116                let ptr = &ptr;
17117                for blk in s..e {
17118                    let (a, b) = (blk * 4096, ((blk + 1) * 4096).min(n));
17119                    for i in a..b {
17120                        unsafe { *ptr.0.add(i) += src[i] };
17121                    }
17122                }
17123            };
17124            p.run_rows(n.div_ceil(4096), &f);
17125        }
17126        _ => {
17127            for (d, v) in dst.iter_mut().zip(&src[..n]) {
17128                *d += *v;
17129            }
17130        }
17131    }
17132}
17133
17134/// `c[m, n] += Σ_t left[t, m]·right[t, n]`, with `left` stored [m, t] and
17135/// `right` [t, n]. Tiled over the rows of `c` so a tile stays in cache
17136/// while each token's `right` row streams past it once, and parallel
17137/// over tiles.
17138fn accum_outer(
17139    c: &mut [f32],
17140    m: usize,
17141    n: usize,
17142    b: usize,
17143    left: &[f32],
17144    right: &[f32],
17145    pool: Option<&Pool>,
17146) {
17147    const TILE: usize = 32;
17148    let tiles = m.div_ceil(TILE);
17149    let cp = SendMut(c.as_mut_ptr());
17150    let body = |ti: usize| {
17151        let cp = &cp;
17152        let i0 = ti * TILE;
17153        let i1 = (i0 + TILE).min(m);
17154        for t in 0..b {
17155            let r = &right[t * n..t * n + n];
17156            for i in i0..i1 {
17157                let a = left[i * b + t];
17158                if a == 0.0 {
17159                    continue;
17160                }
17161                // SAFETY: tiles partition c's rows; workers never overlap.
17162                let row = unsafe { std::slice::from_raw_parts_mut(cp.0.add(i * n), n) };
17163                for (o, v) in row.iter_mut().zip(r) {
17164                    *o += a * *v;
17165                }
17166            }
17167        }
17168    };
17169    match pool {
17170        Some(p) if tiles > 1 => p.run_rows(tiles, &|s, e| {
17171            for ti in s..e {
17172                body(ti);
17173            }
17174        }),
17175        _ => {
17176            for ti in 0..tiles {
17177                body(ti);
17178            }
17179        }
17180    }
17181}
17182
17183/// Write what the calibration accumulated: `gss.<L>.f32` and `ya.<L>.f32`.
17184pub fn refit_flush() -> usize {
17185    let Some((dir, map)) = refit_dir() else {
17186        return 0;
17187    };
17188    let guard = map.lock().unwrap();
17189    let mut n = 0;
17190    for (li, acc) in guard.0.iter() {
17191        // A silently truncated write here is a Gram that reshapes to
17192        // nothing an hour later — say it out loud instead.
17193        let w = |name: &str, v: &[f32]| {
17194            let path = format!("{dir}/{name}.{li}.f32");
17195            let bytes: Vec<u8> = v.iter().flat_map(|x| x.to_le_bytes()).collect();
17196            match std::fs::write(&path, &bytes) {
17197                Ok(()) => {}
17198                Err(e) => eprintln!(
17199                    "refit: FAILED to write {path} ({} MB): {e}",
17200                    bytes.len() / 1_000_000
17201                ),
17202            }
17203        };
17204        w("gss", &acc.gss);
17205        w("ya", &acc.ya);
17206        println!(
17207            "refit L{li}: {} support, {} tokens, hidden {}",
17208            acc.support.len(),
17209            acc.tokens,
17210            acc.hidden
17211        );
17212        n += 1;
17213    }
17214    n
17215}
17216
17217/// `CMF_FFN_ADUMP=<prefix>` — append every probed token's FFN activation
17218/// row to `<prefix>.<layer>.f16`. The co-activation record: which
17219/// neurons fire together, which is what a tube has to group if a token
17220/// is ever going to open one tube instead of sixteen.
17221fn adump_row(li: usize, g: &[f32]) {
17222    use std::io::Write as _;
17223    static FILES: std::sync::OnceLock<
17224        Option<(
17225            String,
17226            std::sync::Mutex<std::collections::HashMap<usize, std::fs::File>>,
17227        )>,
17228    > = std::sync::OnceLock::new();
17229    let Some((prefix, map)) = FILES
17230        .get_or_init(|| {
17231            std::env::var("CMF_FFN_ADUMP")
17232                .ok()
17233                .map(|p| (p, std::sync::Mutex::new(std::collections::HashMap::new())))
17234        })
17235        .as_ref()
17236    else {
17237        return;
17238    };
17239    // `CMF_FFN_ADUMP_FROM/_TO` narrow the dump to a layer span, so a big
17240    // calibration run fits on disk in a few passes instead of one.
17241    static SPAN: std::sync::OnceLock<(usize, usize)> = std::sync::OnceLock::new();
17242    let (from, to) = *SPAN.get_or_init(|| {
17243        let g = |k: &str, d: usize| {
17244            std::env::var(k)
17245                .ok()
17246                .and_then(|v| v.parse().ok())
17247                .unwrap_or(d)
17248        };
17249        (
17250            g("CMF_FFN_ADUMP_FROM", 0),
17251            g("CMF_FFN_ADUMP_TO", usize::MAX),
17252        )
17253    });
17254    if li < from || li > to {
17255        return;
17256    }
17257    let mut map = map.lock().unwrap();
17258    let f = map.entry(li).or_insert_with(|| {
17259        std::fs::File::create(format!("{prefix}.{li}.f16")).expect("adump file")
17260    });
17261    let mut bytes = Vec::with_capacity(g.len() * 2);
17262    for v in g {
17263        bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(*v).to_le_bytes());
17264    }
17265    let _ = f.write_all(&bytes);
17266}
17267
17268/// `CMF_FFN_ORACLE_TOPK` — keep only the k largest |silu(g)·u| of each
17269/// token and zero the rest. Not a serving mode: it is the CEILING of
17270/// contextual sparsity — what a per-token router would be chasing —
17271/// measured by cheating, since the selection reads the very activations
17272/// it would have to predict.
17273fn oracle_topk() -> usize {
17274    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17275    *K.get_or_init(|| {
17276        std::env::var("CMF_FFN_ORACLE_TOPK")
17277            .ok()
17278            .and_then(|v| v.parse().ok())
17279            .unwrap_or(0)
17280    })
17281}
17282
17283/// `CMF_FFN_GATE_TOPK` — the REALIZABLE cousin of the oracle: rank the
17284/// neurons by their gate alone (which the kernel has computed anyway
17285/// before it reads `up`), keep the k best, and drop the rest. Every
17286/// dropped neuron's `up` row and `down` column stay unread, so this is
17287/// the sparsity a serving path can actually take without a router.
17288fn gate_topk() -> usize {
17289    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17290    *K.get_or_init(|| {
17291        std::env::var("CMF_FFN_GATE_TOPK")
17292            .ok()
17293            .and_then(|v| v.parse().ok())
17294            .unwrap_or(0)
17295    })
17296}
17297
17298/// `CMF_FFN_GATE_BLOCK` — select in blocks of B neurons instead of one
17299/// by one. A scattered per-neuron choice cannot be read efficiently (a
17300/// row at a time, no prefetch runway); a block of 32 is a contiguous
17301/// 32-row slab of `up` and of the transposed `down`, which the ordinary
17302/// kernels stream. The question the measurement answers is what the
17303/// block costs in quality.
17304fn gate_block() -> usize {
17305    static B: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17306    *B.get_or_init(|| {
17307        std::env::var("CMF_FFN_GATE_BLOCK")
17308            .ok()
17309            .and_then(|v| v.parse().ok())
17310            .unwrap_or(1)
17311    })
17312}
17313
17314/// Zero all but the `k` largest BLOCKS (by summed square) of a row.
17315fn keep_top_blocks(g: &mut [f32], keep_n: usize, block: usize) {
17316    let n = g.len();
17317    let nb = n.div_ceil(block);
17318    let kb = (keep_n.div_ceil(block)).clamp(1, nb);
17319    if kb >= nb {
17320        return;
17321    }
17322    let mut score: Vec<f32> = (0..nb)
17323        .map(|b| {
17324            g[b * block..((b + 1) * block).min(n)]
17325                .iter()
17326                .map(|v| v * v)
17327                .sum::<f32>()
17328        })
17329        .collect();
17330    let mut ord = score.clone();
17331    let (_, kth, _) = ord.select_nth_unstable_by(kb - 1, |a, b| {
17332        b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17333    });
17334    let thr = *kth;
17335    for b in 0..nb {
17336        if score[b] < thr {
17337            g[b * block..((b + 1) * block).min(n)].fill(0.0);
17338        }
17339    }
17340    score.clear();
17341}
17342
17343/// Zero all but the `k` largest magnitudes of one token's activation row.
17344fn keep_top_k(g: &mut [f32], k: usize) {
17345    if gate_block() > 1 {
17346        return keep_top_blocks(g, k, gate_block());
17347    }
17348    let n = g.len();
17349    if k == 0 || k >= n {
17350        return;
17351    }
17352    let mut mag: Vec<f32> = g.iter().map(|v| v.abs()).collect();
17353    let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
17354        b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17355    });
17356    let thr = *kth;
17357    for v in g.iter_mut() {
17358        if v.abs() < thr {
17359            *v = 0.0;
17360        }
17361    }
17362}
17363
17364/// `CMF_FFN_PROBE_SQ` — accumulate Σa², so the dump divided by the token
17365/// count and square-rooted is the RMS activation trace Patent 12 weights
17366/// its matrices by.
17367fn probe_sq() -> bool {
17368    static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17369    *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SQ").is_ok())
17370}
17371
17372/// `CMF_FFN_PROBE_SIGNED` — accumulate the SIGNED activation sum
17373/// instead of its magnitude: what a dropped neuron contributes ON
17374/// AVERAGE, which is the bias a narrowed FFN can add back for free.
17375fn probe_signed() -> bool {
17376    static S: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
17377    *S.get_or_init(|| std::env::var("CMF_FFN_PROBE_SIGNED").is_ok())
17378}
17379
17380/// `CMF_FFN_MEANFILL=<file>` — a masked-out neuron contributes its MEAN
17381/// activation instead of zero (`u32 layers, u32 inter, f32[…]`, the mass
17382/// dump layout, holding per-neuron means). Dropping a neuron outright
17383/// also drops its average contribution, which shifts the layer output by
17384/// a constant; filling the mean back is one add per layer and costs no
17385/// bytes off the bus. This is the measurement arm — in a tube file the
17386/// same correction ships as a per-task bias vector.
17387fn meanfill() -> Option<&'static (usize, Vec<f32>)> {
17388    static M: std::sync::OnceLock<Option<(usize, Vec<f32>)>> = std::sync::OnceLock::new();
17389    M.get_or_init(|| {
17390        let p = std::env::var("CMF_FFN_MEANFILL").ok()?;
17391        let b = std::fs::read(&p).ok()?;
17392        let inter = u32::from_le_bytes(b[4..8].try_into().ok()?) as usize;
17393        let vals: Vec<f32> = b[8..]
17394            .chunks_exact(4)
17395            .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
17396            .collect();
17397        eprintln!("meanfill: {} value(s), inter {inter}", vals.len());
17398        Some((inter, vals))
17399    })
17400    .as_ref()
17401}
17402
17403/// `CMF_FFN_PROBE_TOPK` — 0 (default) = accumulate mass, k>0 = count
17404/// how often a neuron lands in a token's top k.
17405fn probe_topk() -> usize {
17406    static K: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
17407    *K.get_or_init(|| {
17408        std::env::var("CMF_FFN_PROBE_TOPK")
17409            .ok()
17410            .and_then(|v| v.parse().ok())
17411            .unwrap_or(0)
17412    })
17413}
17414
17415thread_local! {
17416    /// DTG-MA activation probe: per-layer per-neuron Σ|silu(g)·u|
17417    /// accumulator, alive only during `Pipeline::probe_ffn_mass`.
17418    static FFN_PROBE: std::cell::RefCell<Option<Vec<Vec<f64>>>> =
17419        const { std::cell::RefCell::new(None) };
17420}
17421
17422/// Per-token structured sparsity, paid for in bytes.
17423///
17424/// The gate is the cheapest third of an FFN and it already says which
17425/// neurons matter: `silu(gate)` near zero means the neuron contributes
17426/// nothing whatever `up` says. So compute every gate, keep the `k`
17427/// loudest, and read ONLY those neurons' `up` rows and `down` rows —
17428/// the latter needs `down_proj` stored transposed, otherwise a neuron's
17429/// down weights are a strided column and "reading only those" costs a
17430/// full cache line each.
17431///
17432/// Returns `None` when the file has no transposed `down` (the caller
17433/// then runs the ordinary dense path).
17434fn dense_ffn_dynamic(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, k: usize) -> Option<Vec<f32>> {
17435    // The scatter path reads individual rows/columns and cannot express the
17436    // per-matrix signed FWHT boundary.  Let the descriptor-aware dense path
17437    // handle Prism files rather than silently running an unrotated sparse
17438    // approximation.
17439    if d.gate_proj.has_prism_contract()
17440        || d.up_proj.has_prism_contract()
17441        || d.down_proj.has_prism_contract()
17442    {
17443        return None;
17444    }
17445    let dt = d.down_t.as_ref()?;
17446    let inter = d.gate_proj.rows();
17447    let hidden = dt.cols();
17448    if k == 0 || k >= inter || d.act != Act::Silu {
17449        return None;
17450    }
17451    DYN_SCRATCH.with(|sc| {
17452        let mut sc = sc.borrow_mut();
17453        let DynScratch {
17454            g,
17455            mag,
17456            live,
17457            parts,
17458        } = &mut *sc;
17459        g.resize(inter, 0.0);
17460        d.gate_proj.matvec(x, g, pool);
17461        for v in g.iter_mut() {
17462            *v = inference::silu(*v);
17463        }
17464        // The k-th largest |silu(gate)| is the threshold; ties keep more,
17465        // which is the safe side.
17466        mag.clear();
17467        mag.extend(g.iter().map(|v| v.abs()));
17468        let (_, kth, _) = mag.select_nth_unstable_by(k - 1, |a, b| {
17469            b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal)
17470        });
17471        let thr = *kth;
17472        live.clear();
17473        live.extend((0..inter as u32).filter(|&n| g[n as usize].abs() >= thr));
17474        let mut out = vec![0.0f32; hidden];
17475        match pool {
17476            Some(p) if live.len() >= 64 => {
17477                let nw = p.n_workers() + 1;
17478                parts.clear();
17479                parts.resize(nw * hidden, 0.0);
17480                let ptr = SendMut(parts.as_mut_ptr());
17481                let n = live.len();
17482                let live_ref: &[u32] = live;
17483                let g_ref: &[f32] = g;
17484                p.run(&|w, workers| {
17485                    let chunk = n.div_ceil(workers);
17486                    let (s, e) = (w * chunk, ((w + 1) * chunk).min(n));
17487                    if s >= e {
17488                        return;
17489                    }
17490                    WORKER_SCRATCH.with(|ws| {
17491                        let mut ws = ws.borrow_mut();
17492                        let [scratch, acc] = &mut *ws;
17493                        scratch.resize(hidden.max(x.len()), 0.0);
17494                        acc.clear();
17495                        acc.resize(hidden, 0.0);
17496                        for (o, &nrm) in live_ref[s..e].iter().enumerate() {
17497                            // One neuron of runway: the next row's lines
17498                            // start moving while this one is multiplied.
17499                            if let Some(&nx) = live_ref[s..e].get(o + 1) {
17500                                d.up_proj.prefetch_row(nx as usize);
17501                                dt.prefetch_row(nx as usize);
17502                            }
17503                            let idx = nrm as usize;
17504                            let up = d.up_proj.row_dot(idx, x, scratch);
17505                            let a = g_ref[idx] * up;
17506                            if a != 0.0 {
17507                                dt.add_row_scaled(idx, a, acc, scratch);
17508                            }
17509                        }
17510                        for (j, v) in acc.iter().enumerate() {
17511                            unsafe { *ptr.at(w * hidden + j) = *v };
17512                        }
17513                    });
17514                });
17515                for w in 0..nw {
17516                    for (j, o) in out.iter_mut().enumerate() {
17517                        *o += parts[w * hidden + j];
17518                    }
17519                }
17520            }
17521            _ => {
17522                WORKER_SCRATCH.with(|ws| {
17523                    let mut ws = ws.borrow_mut();
17524                    let [scratch, _acc] = &mut *ws;
17525                    scratch.resize(hidden.max(x.len()), 0.0);
17526                    for &nrm in live.iter() {
17527                        let idx = nrm as usize;
17528                        let up = d.up_proj.row_dot(idx, x, scratch);
17529                        let a = g[idx] * up;
17530                        if a != 0.0 {
17531                            dt.add_row_scaled(idx, a, &mut out, scratch);
17532                        }
17533                    }
17534                });
17535            }
17536        }
17537        Some(out)
17538    })
17539}
17540
17541/// Caller-side scratch of the dynamic path — one allocation per thread,
17542/// not one per layer per token (that alone cost a third of the decode).
17543struct DynScratch {
17544    g: Vec<f32>,
17545    mag: Vec<f32>,
17546    live: Vec<u32>,
17547    parts: Vec<f32>,
17548}
17549
17550thread_local! {
17551    static DYN_SCRATCH: std::cell::RefCell<DynScratch> = const {
17552        std::cell::RefCell::new(DynScratch {
17553            g: Vec::new(),
17554            mag: Vec::new(),
17555            live: Vec::new(),
17556            parts: Vec::new(),
17557        })
17558    };
17559    /// Pool-worker scratch: the row buffer and this worker's partial sum.
17560    static WORKER_SCRATCH: std::cell::RefCell<[Vec<f32>; 2]> =
17561        const { std::cell::RefCell::new([Vec::new(), Vec::new()]) };
17562}
17563
17564/// `dense_ffn_cpu` with a per-visit mask landing on the activations —
17565/// the masked-inference fast path's decode arm. Full fused quant
17566/// compute, closed neurons zeroed before down: arithmetically the
17567/// pruned network, no dequant, no weight bytes touched.
17568fn dense_ffn_masked(d: &DenseFfn, x: &[f32], pool: Option<&Pool>, mask_row: &[u8]) -> Vec<f32> {
17569    let inter = d.gate_proj.rows();
17570    FFN_SCRATCH.with(|s| {
17571        let mut s = s.borrow_mut();
17572        let [g, u, ..] = &mut *s;
17573        g.resize(inter, 0.0);
17574        if d.act == Act::Silu && QTensor::matvec_silu_mul(&d.gate_proj, &d.up_proj, x, g, pool) {
17575            // g holds silu(gate)·up.
17576        } else {
17577            u.resize(inter, 0.0);
17578            QTensor::matvec_many([&d.gate_proj, &d.up_proj], x, [g, u], pool);
17579            for i in 0..inter {
17580                g[i] = d.act.combine(g[i], u[i]);
17581            }
17582        }
17583        zero_masked_cols(g, 1, inter, mask_row);
17584        let mut out = attention::take_buf(d.down_proj.rows());
17585        d.down_proj.matvec(g, &mut out, pool);
17586        out
17587    })
17588}
17589
17590/// Dense FFN as one GPU submission via the MoE block path (single
17591/// expert, weight 1.0): gate → silu·up → down chained in one command
17592/// buffer, intermediate activations device-resident. None → weights
17593/// not q8-mapped in the primary shard / over the VRAM budget / backend
17594/// refusal → honest CPU path.
17595fn dense_ffn_gpu(d: &DenseFfn, x: &[f32], _pool: Option<&Pool>) -> Option<Vec<f32>> {
17596    if d.gate_proj.has_prism_contract()
17597        || d.up_proj.has_prism_contract()
17598        || d.down_proj.has_prism_contract()
17599    {
17600        return None;
17601    }
17602    // The GPU block hardcodes SiLU; GeLU FFNs (Gemma) stay on CPU.
17603    if d.act != Act::Silu {
17604        return None;
17605    }
17606    // Threshold: tiny FFNs are not worth a submission (q1 excepted —
17607    // see the caller's gate).
17608    if d.gate_proj.rows() < crate::gpu::min_rows() && !d.gate_proj.is_q1() {
17609        return None;
17610    }
17611    let mut jobs: Vec<crate::gpu::MoeJob> = Vec::with_capacity(1);
17612    let mut model_ref = None;
17613    moe_push_job(d, x, 1.0, &mut jobs, &mut model_ref)?;
17614    let model = model_ref?;
17615    let hidden = jobs[0].down.1;
17616    let mut out = attention::take_buf(hidden);
17617    if crate::gpu::moe_block(&model, &jobs, &mut out) {
17618        Some(out)
17619    } else {
17620        let mut out = out;
17621        attention::recycle_buf(&mut out);
17622        None
17623    }
17624}
17625
17626/// q8-mapped primary-shard tensor parts for a GPU job: q8_2f carries
17627/// its column field, q8_row runs with empty col slices (the backend
17628/// skips the multiply). Shared by the MoE block and the dense-FFN
17629/// single-job path.
17630#[allow(clippy::type_complexity)]
17631#[allow(clippy::type_complexity)]
17632pub(crate) fn moe_parts(
17633    t: &QTensor,
17634) -> Option<(
17635    &std::sync::Arc<cortiq_core::CmfModel>,
17636    usize,
17637    usize,
17638    usize,
17639    &[f32],
17640    &[f32],
17641    bool,
17642    bool,
17643    bool,
17644)> {
17645    match t {
17646        QTensor::Mapped {
17647            model,
17648            idx,
17649            dtype: dt @ (cortiq_core::TensorDtype::Q8_2f | cortiq_core::TensorDtype::Q8Row),
17650            rows,
17651            cols,
17652            row_scale,
17653            col_field,
17654            ..
17655        } if (*dt == cortiq_core::TensorDtype::Q8Row) || !col_field.is_empty() => Some((
17656            model, *idx, *rows, *cols, row_scale, col_field, false, false, false,
17657        )),
17658        // q1: tile-embedded scales — empty rs/col slices, raw xs.
17659        QTensor::Mapped {
17660            model,
17661            idx,
17662            dtype: cortiq_core::TensorDtype::Q1,
17663            rows,
17664            cols,
17665            ..
17666        } => Some((
17667            model,
17668            *idx,
17669            *rows,
17670            *cols,
17671            &[][..],
17672            &[][..],
17673            true,
17674            false,
17675            false,
17676        )),
17677        // q4_tiled: 18-byte tiles with embedded f16 scales — raw xs.
17678        QTensor::Mapped {
17679            model,
17680            idx,
17681            dtype: cortiq_core::TensorDtype::Q4Tiled,
17682            rows,
17683            cols,
17684            ..
17685        } => Some((
17686            model,
17687            *idx,
17688            *rows,
17689            *cols,
17690            &[][..],
17691            &[][..],
17692            false,
17693            true,
17694            false,
17695        )),
17696        // q4tp: same raw-xs contract, different stride and scale plane.
17697        QTensor::Mapped {
17698            model,
17699            idx,
17700            dtype: cortiq_core::TensorDtype::Q4TiledP,
17701            rows,
17702            cols,
17703            ..
17704        } => Some((
17705            model,
17706            *idx,
17707            *rows,
17708            *cols,
17709            &[][..],
17710            &[][..],
17711            false,
17712            true,
17713            false,
17714        )),
17715        // q2tp: the 2-bit expert plane of the mixed profile — q4 family
17716        // for stride bookkeeping, flagged q2 so the trio validation can
17717        // demand a q4tp down.
17718        QTensor::Mapped {
17719            model,
17720            idx,
17721            dtype: cortiq_core::TensorDtype::Q2TiledP,
17722            rows,
17723            cols,
17724            ..
17725        } => Some((
17726            model,
17727            *idx,
17728            *rows,
17729            *cols,
17730            &[][..],
17731            &[][..],
17732            false,
17733            true,
17734            true,
17735        )),
17736        _ => None,
17737    }
17738}
17739
17740/// Map a MoE onto the Metal token graph's contract: f32 router, a
17741/// shared expert (gated — Qwen — or ungated at weight 1 — DeepSeek-V3 /
17742/// HunYuan hy_v3), softmax or sigmoid scores with an optional selection
17743/// bias and routed scale, experts uniformly q4tp (or the mixed profile:
17744/// q2tp gate/up over a q4tp down). τ routers, masks, per-expert scales
17745/// and Gemma's router-input norm refuse here — those semantics stay on
17746/// the CPU path.
17747#[cfg(target_os = "macos")]
17748fn metal_moe_graph_parts(m: &MoeFfn, hidden: usize) -> Option<crate::gpu::GpuMoe<'_>> {
17749    if m.router_input_norm
17750        || m.route_tau.is_some()
17751        || m.mask.is_some()
17752        || m.per_expert_scale.is_some()
17753        || m.experts.is_empty()
17754        || m.top_k == 0
17755        || m.resonance.is_some()
17756    {
17757        return None;
17758    }
17759    // The select kernel always fills the shared slot: a model without a
17760    // shared expert (LFM2-MoE) stays on the CPU path here.
17761    let (sh, sg) = match &m.shared {
17762        Some((sh, sg)) => (sh, sg.as_ref()),
17763        None => return None,
17764    };
17765    let (rf, rr, rc) = m.router.f32_parts()?;
17766    if rr != m.experts.len() || rc != hidden {
17767        return None;
17768    }
17769    let shared_gated = sg.is_some();
17770    let sf = match sg {
17771        Some(sg) => {
17772            let (sf, sr, sc) = sg.f32_parts()?;
17773            if sr * sc != hidden {
17774                return None;
17775            }
17776            sf
17777        }
17778        // Ungated: the router's first row stands in for the gate matvec
17779        // (its logit is never read — the kernel pins weight 1).
17780        None => &rf[..hidden],
17781    };
17782    if let Some(b) = &m.expert_bias {
17783        if b.len() != m.experts.len() {
17784            return None;
17785        }
17786    }
17787    let inter = m.experts[0].gate_proj.rows();
17788    // The first expert's gate decides the profile; every trio (shared
17789    // included) must agree — the jobs ladder flips ONE kernel for all.
17790    let gu_q2 = m.experts[0].gate_proj.mapped_q2tp().is_some();
17791    let trio = |e: &DenseFfn| -> Option<(usize, usize, usize)> {
17792        if e.act != Act::Silu
17793            || e.gate_proj.rows() != inter
17794            || e.gate_proj.cols() != hidden
17795            || e.up_proj.rows() != inter
17796            || e.up_proj.cols() != hidden
17797            || e.down_proj.rows() != hidden
17798            || e.down_proj.cols() != inter
17799        {
17800            return None;
17801        }
17802        let pick = |t: &QTensor| -> Option<usize> {
17803            if gu_q2 {
17804                t.mapped_q2tp().map(|(_, i)| i)
17805            } else {
17806                t.mapped_q4tp().map(|(_, i)| i)
17807            }
17808        };
17809        Some((
17810            pick(&e.gate_proj)?,
17811            pick(&e.up_proj)?,
17812            e.down_proj.mapped_q4tp().map(|(_, i)| i)?,
17813        ))
17814    };
17815    let experts = m.experts.iter().map(trio).collect::<Option<Vec<_>>>()?;
17816    let shared = trio(sh)?;
17817    Some(crate::gpu::GpuMoe {
17818        router: rf,
17819        sgate: sf,
17820        experts,
17821        shared,
17822        n_exp: m.experts.len(),
17823        top_k: m.top_k,
17824        inter,
17825        norm_topk: m.norm_topk_prob,
17826        route_scale: m.routed_scaling,
17827        gu_q2,
17828        sigmoid: m.router_sigmoid,
17829        bias: m.expert_bias.as_deref(),
17830        shared_gated,
17831    })
17832}
17833
17834/// Build one gate/up/down GPU job from three tensors. `moe_push_job` is the
17835/// DenseFfn-shaped caller; architectures that keep their experts in their own
17836/// structs (DeepSeek-V4) come here directly.
17837pub(crate) fn moe_push_job_parts<'a>(
17838    gate: &'a QTensor,
17839    up: &'a QTensor,
17840    down: &'a QTensor,
17841    x: &[f32],
17842    w: f32,
17843    swiglu_limit: f32,
17844    jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17845    model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17846) -> Option<()> {
17847    use crate::qtensor::prescale;
17848    let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(gate)?;
17849    let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(up)?;
17850    let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(down)?;
17851    if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17852        return None; // mixed-dtype trio — honest CPU path
17853    }
17854    // The 2-bit profile is gate/up q2tp over a PLAIN q4tp down; any other
17855    // 2-bit arrangement stays on the CPU.
17856    if gq2 && (dq2 || !dq4 || down.mapped_q4tp().is_none()) {
17857        return None;
17858    }
17859    if !gq2 && dq2 {
17860        return None;
17861    }
17862    model_ref.get_or_insert_with(|| gm.clone());
17863    let dt = |cf: &[f32]| {
17864        if cf.is_empty() {
17865            cortiq_core::TensorDtype::Q8Row
17866        } else {
17867            cortiq_core::TensorDtype::Q8_2f
17868        }
17869    };
17870    jobs.push(crate::gpu::MoeJob {
17871        gate: (gi, gr, gc, grs),
17872        up: (ui, ur, uc, urs),
17873        down: (di, dr, dc, drs),
17874        xs_gate: prescale(x, gcf, dt(gcf)).into_owned(),
17875        xs_up: prescale(x, ucf, dt(ucf)).into_owned(),
17876        down_col: dcf,
17877        w,
17878        q1: gq1,
17879        q4t: gq4 && !gq2 && gate.mapped_q4tp().is_none(),
17880        q4tp: gq4 && (gq2 || gate.mapped_q4tp().is_some()),
17881        gu_q2: gq2,
17882        swiglu_limit,
17883    });
17884    Some(())
17885}
17886
17887/// Build one gate/up/down GPU job (see `moe_parts`).
17888fn moe_push_job<'a>(
17889    d: &'a DenseFfn,
17890    x: &[f32],
17891    w: f32,
17892    jobs: &mut Vec<crate::gpu::MoeJob<'a>>,
17893    model_ref: &mut Option<std::sync::Arc<cortiq_core::CmfModel>>,
17894) -> Option<()> {
17895    use crate::qtensor::prescale;
17896    if d.act != Act::Silu {
17897        return None; // GPU block hardcodes SiLU
17898    }
17899    let (gm, gi, gr, gc, grs, gcf, gq1, gq4, gq2) = moe_parts(&d.gate_proj)?;
17900    let (_, ui, ur, uc, urs, ucf, uq1, uq4, uq2) = moe_parts(&d.up_proj)?;
17901    let (_, di, dr, dc, drs, dcf, dq1, dq4, dq2) = moe_parts(&d.down_proj)?;
17902    if gq1 != uq1 || uq1 != dq1 || gq4 != uq4 || uq4 != dq4 || gq2 != uq2 {
17903        return None; // mixed-dtype trio — honest CPU path
17904    }
17905    if gq2 && (dq2 || !dq4 || d.down_proj.mapped_q4tp().is_none()) {
17906        return None;
17907    }
17908    if !gq2 && dq2 {
17909        return None;
17910    }
17911    model_ref.get_or_insert_with(|| gm.clone());
17912    let gdt = if gcf.is_empty() {
17913        cortiq_core::TensorDtype::Q8Row
17914    } else {
17915        cortiq_core::TensorDtype::Q8_2f
17916    };
17917    let udt = if ucf.is_empty() {
17918        cortiq_core::TensorDtype::Q8Row
17919    } else {
17920        cortiq_core::TensorDtype::Q8_2f
17921    };
17922    jobs.push(crate::gpu::MoeJob {
17923        gate: (gi, gr, gc, grs),
17924        up: (ui, ur, uc, urs),
17925        down: (di, dr, dc, drs),
17926        xs_gate: prescale(x, gcf, gdt).into_owned(),
17927        xs_up: prescale(x, ucf, udt).into_owned(),
17928        down_col: dcf,
17929        w,
17930        q1: gq1,
17931        q4t: gq4 && !gq2 && d.gate_proj.mapped_q4tp().is_none(),
17932        q4tp: gq4 && (gq2 || d.gate_proj.mapped_q4tp().is_some()),
17933        gu_q2: gq2,
17934        swiglu_limit: 0.0,
17935    });
17936    Some(())
17937}
17938
17939/// Sparse dense-FFN directly on QUANTIZED weights (mask × mmap): reads
17940/// ONLY the active neurons' gate/up rows and down columns from the mmap
17941/// — no full-matrix dequant, no f32 model copy. This is what lets a
17942/// masked big model run at quantized RSS (the historical mask path
17943/// forced the whole model to f32). Semantics identical to the f32
17944/// sparse path within quant tolerance.
17945fn sparse_ffn_quant(
17946    d: &DenseFfn,
17947    x: &[f32],
17948    active: &[u16],
17949    hidden: usize,
17950    pool: Option<&Pool>,
17951) -> Vec<f32> {
17952    let n = active.len();
17953    let inter = d.gate_proj.rows();
17954    let mut act = vec![0.0f32; n];
17955    // Scratch is needed if EITHER projection is group-packed (q4/vbit);
17956    // gate/up normally share a dtype but sizing on both is robust.
17957    let need_scratch = !(d.gate_proj.sparse_col_ok() && d.up_proj.sparse_col_ok());
17958    let compute = |ai: usize| -> f32 {
17959        let idx = active[ai] as usize;
17960        if idx >= inter {
17961            return 0.0; // defensive parity with the f32 sparse path
17962        }
17963        let mut s = if need_scratch {
17964            vec![0.0f32; hidden]
17965        } else {
17966            Vec::new()
17967        };
17968        let gate = d.gate_proj.row_dot(idx, x, &mut s);
17969        let up = d.up_proj.row_dot(idx, x, &mut s);
17970        d.act.combine(gate, up)
17971    };
17972    match pool {
17973        Some(p) if n >= 256 => {
17974            let ptr = SendMut(act.as_mut_ptr());
17975            p.run(&|widx, nw| {
17976                let chunk = n.div_ceil(nw);
17977                let (s, e) = (widx * chunk, ((widx + 1) * chunk).min(n));
17978                for ai in s..e {
17979                    unsafe { *ptr.at(ai) = compute(ai) };
17980                }
17981            });
17982        }
17983        _ => {
17984            for (ai, a) in act.iter_mut().enumerate() {
17985                *a = compute(ai);
17986            }
17987        }
17988    }
17989    // Scatter through active down columns (reads only those columns).
17990    let mut out = vec![0.0f32; hidden];
17991    for (ai, &idx) in active.iter().enumerate() {
17992        let w = act[ai];
17993        if w.abs() >= 1e-12 && (idx as usize) < inter {
17994            d.down_proj.add_col_scaled(idx as usize, w, &mut out);
17995        }
17996    }
17997    out
17998}
17999
18000/// Test-only re-export of the private sparse-quant FFN (mask × mmap gate).
18001#[doc(hidden)]
18002pub fn sparse_ffn_quant_for_test(
18003    d: &DenseFfn,
18004    x: &[f32],
18005    active: &[u16],
18006    hidden: usize,
18007) -> Vec<f32> {
18008    sparse_ffn_quant(d, x, active, hidden, None)
18009}
18010
18011/// Dequantize a DenseFfn's three matrices to f32 (transient; only the
18012/// q4/vbit-masked fallback uses it — the memory-lean path is
18013/// sparse_ffn_quant). Reuses row_f32 row-by-row.
18014fn dequant_dense_f32(d: &DenseFfn) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
18015    let deq = |t: &QTensor| -> Vec<f32> {
18016        let (rows, cols) = (t.rows(), t.cols());
18017        let mut out = vec![0.0f32; rows * cols];
18018        for r in 0..rows {
18019            t.row_f32(r, &mut out[r * cols..(r + 1) * cols]);
18020        }
18021        out
18022    };
18023    (deq(&d.gate_proj), deq(&d.up_proj), deq(&d.down_proj))
18024}
18025
18026/// Pointer wrapper for the worker-pool scatter (same pattern as qtensor).
18027struct SendMut(*mut f32);
18028unsafe impl Send for SendMut {}
18029unsafe impl Sync for SendMut {}
18030impl SendMut {
18031    #[inline]
18032    // Deliberate unsynchronized scatter: pool workers write disjoint indices
18033    // in parallel, so returning `&mut` from `&self` is intentional here.
18034    #[allow(clippy::mut_from_ref)]
18035    unsafe fn at(&self, i: usize) -> &mut f32 {
18036        unsafe { &mut *self.0.add(i) }
18037    }
18038}
18039
18040/// Router → (selected experts in torch.topk order, per-expert score
18041/// vector, normalizer). The final weight of expert `e` is `p[e] / wsum`.
18042///
18043/// Two regimes share this. Qwen: softmax over ALL experts, top-k of the
18044/// probabilities, optional renorm — `router_sigmoid=false`, no bias,
18045/// scale 1 → bit-identical to the historical path. LFM2-MoE /
18046/// DeepSeek-V3 `noaux_tc`: per-expert sigmoid scores, an optional
18047/// selection bias (top-k CHOICE only; weights stay unbiased), a 1e-6 renorm
18048/// floor and a routed scale. Architectures whose reference uses a different
18049/// sigmoid denominator floor (for example GLM-5's `1e-20`) call
18050/// [`moe_route_with_eps`] directly; the historical generic path remains
18051/// unchanged.
18052pub(crate) fn moe_route(
18053    logits: &[f32],
18054    m: &MoeFfn,
18055    allowed: Option<&[bool]>,
18056) -> (Vec<usize>, Vec<f32>, f32) {
18057    moe_route_with_eps(logits, m, allowed, 1e-6)
18058}
18059
18060/// Router implementation with an explicit sigmoid renormalization floor.
18061///
18062/// GLM-5.3's source computes `sum(selected_scores) + 1e-20`; using the
18063/// generic 1e-6 floor there is not a harmless tolerance difference when all
18064/// logits are very negative: it collapses the routed branch toward zero
18065/// instead of normalizing the selected experts. Keeping the epsilon parameter
18066/// here avoids changing the established Qwen/LFM2 contract while allowing
18067/// each architecture to preserve its own numerical semantics.
18068pub(crate) fn moe_route_with_eps(
18069    logits: &[f32],
18070    m: &MoeFfn,
18071    allowed: Option<&[bool]>,
18072    sigmoid_denom_eps: f32,
18073) -> (Vec<usize>, Vec<f32>, f32) {
18074    let ne = logits.len();
18075    // Expert restriction: the static env mask (CMF_MOE_MASK) AND the
18076    // active task mask's expert fields (spec §5) both narrow the
18077    // candidate set; selection happens over the admitted experts only.
18078    // With norm_topk the kept weights renormalize below; without it
18079    // the excluded mass is honestly dropped.
18080    let admit = |e: usize| {
18081        m.mask.as_ref().is_none_or(|mk| mk[e])
18082            && allowed.is_none_or(|a| a.get(e).copied().unwrap_or(false))
18083    };
18084    // The resonance router (spec §9.5.1) selects by the RAW score: the
18085    // trainer (`resonance_winner`) and the resident graph
18086    // (`embryo_core_route_pick`) take the first maximum of the scores
18087    // and run the winner with weight 1.0. Selecting through the softmax
18088    // instead is not the same decision: `exp(l − max)` rounds two scores
18089    // closer than 2^-25 (possible below |score| 0.25) to the same 1.0, and
18090    // the lower index would take a token whose score is strictly smaller
18091    // — the trainer's trace and the graph would disagree with this path.
18092    // `−∞` (outside the shell) never wins; with no finite admitted expert
18093    // the generic path below degrades to uniform. The winner's
18094    // probability is 1.0 by construction (a one-hot `p`), so its
18095    // renormalized weight is `routed_scaling` on both norm_topk settings.
18096    if m.resonance.is_some() && m.top_k == 1 {
18097        let mut best: Option<usize> = None;
18098        for e in (0..ne).filter(|&e| admit(e)) {
18099            let l = logits[e];
18100            if l == f32::NEG_INFINITY || l.is_nan() {
18101                continue;
18102            }
18103            if best.is_none_or(|b| l > logits[b]) {
18104                best = Some(e);
18105            }
18106        }
18107        if let Some(b) = best {
18108            let mut p = vec![0.0f32; ne];
18109            p[b] = 1.0;
18110            return (vec![b], p, 1.0 / m.routed_scaling);
18111        }
18112    }
18113    // A `−∞` logit (a grown expert outside its shell, `Resonance::scores`)
18114    // takes probability 0 on both paths: sigmoid(−∞) = 0, exp(−∞ − max) = 0
18115    // — top-1 is the best FINITE expert, its renormalized weight exactly
18116    // 1.0. Every expert at −∞ cannot happen (trunk experts have no shell);
18117    // should it, the softmax would be NaN, so it degrades to uniform.
18118    let p: Vec<f32> = if m.router_sigmoid {
18119        logits.iter().map(|&l| 1.0 / (1.0 + (-l).exp())).collect()
18120    } else {
18121        let mx = logits.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
18122        if mx == f32::NEG_INFINITY {
18123            vec![1.0 / ne.max(1) as f32; ne]
18124        } else {
18125            let mut e: Vec<f32> = logits.iter().map(|&l| (l - mx).exp()).collect();
18126            let s: f32 = e.iter().sum();
18127            for v in &mut e {
18128                *v /= s;
18129            }
18130            e
18131        }
18132    };
18133    let mut idx: Vec<usize> = (0..ne).filter(|&e| admit(e)).collect();
18134    // Descending by selection score, lower index wins ties (torch.topk).
18135    match &m.expert_bias {
18136        Some(b) => idx.sort_unstable_by(|&x, &y| {
18137            (p[y] + b[y])
18138                .partial_cmp(&(p[x] + b[x]))
18139                .unwrap()
18140                .then(x.cmp(&y))
18141        }),
18142        None => idx.sort_unstable_by(|&x, &y| p[y].partial_cmp(&p[x]).unwrap().then(x.cmp(&y))),
18143    }
18144    idx.truncate(m.top_k);
18145    // Adaptive τ-routing: trim the tail experts once the kept mass is
18146    // enough. wsum below renormalizes over the KEPT set, so the output
18147    // stays a proper weighted average.
18148    if let Some(tau) = m.route_tau {
18149        let total: f32 = idx.iter().map(|&e| p[e]).sum();
18150        if total > 0.0 {
18151            let mut acc = 0.0f32;
18152            let mut keep = idx.len();
18153            for (i, &e) in idx.iter().enumerate() {
18154                acc += p[e];
18155                if acc >= tau * total {
18156                    keep = i + 1;
18157                    break;
18158                }
18159            }
18160            idx.truncate(keep);
18161        }
18162    }
18163    let wsum: f32 = if m.norm_topk_prob {
18164        let s: f32 = idx.iter().map(|&e| p[e]).sum();
18165        // Sigmoid routers use their architecture's reference floor; the
18166        // softmax path's probs already sum near 1, so it stays exactly as
18167        // before.
18168        (if m.router_sigmoid {
18169            s + sigmoid_denom_eps
18170        } else {
18171            s
18172        }) / m.routed_scaling
18173    } else {
18174        1.0 / m.routed_scaling
18175    };
18176    (idx, p, wsum)
18177}
18178
18179/// See the call site: one `layer:e1,e2,…` line per routed token.
18180fn moe_trace(idx: &[usize]) {
18181    moe_trace_at(crate::gpu::cur_layer() as i32, idx)
18182}
18183
18184/// The same, for callers that know their layer (DSV4 owns its layers and
18185/// never sets the pipeline's current-layer marker).
18186pub(crate) fn moe_trace_at(li: i32, idx: &[usize]) {
18187    use std::io::Write;
18188    static F: std::sync::OnceLock<Option<std::sync::Mutex<std::fs::File>>> =
18189        std::sync::OnceLock::new();
18190    let Some(f) = F.get_or_init(|| {
18191        let p = std::env::var("CMF_MOE_TRACE").ok()?;
18192        Some(std::sync::Mutex::new(
18193            std::fs::OpenOptions::new()
18194                .create(true)
18195                .append(true)
18196                .open(p)
18197                .ok()?,
18198        ))
18199    }) else {
18200        return;
18201    };
18202    let ids: Vec<String> = idx.iter().map(|e| e.to_string()).collect();
18203    let _ = writeln!(f.lock().unwrap(), "{li}:{}", ids.join(","));
18204}
18205
18206/// MoE FFN: router → top-k experts (see `moe_route`). Only selected
18207/// experts' pages are touched in mmap.
18208pub(crate) fn moe_ffn(
18209    m: &MoeFfn,
18210    x: &[f32],
18211    pool: Option<&Pool>,
18212    allowed: Option<&[bool]>,
18213) -> Vec<f32> {
18214    let r = moe_ffn_route(m, x, pool, allowed);
18215    moe_ffn_experts(m, x, &r, pool)
18216}
18217
18218/// One token's host route through a MoE layer: the chosen experts in
18219/// selection order, the per-expert scores and the normalizer (see
18220/// `moe_route`), plus the raw router logits.
18221pub(crate) struct MoeRoute {
18222    pub idx: Vec<usize>,
18223    pub p: Vec<f32>,
18224    pub wsum: f32,
18225    pub logits: Vec<f32>,
18226}
18227
18228/// The routing half of `moe_ffn`, shared by every executor of the chosen
18229/// experts (the host/per-op path below and the MiMo dynamic device cache,
18230/// `crate::mimo_moe`): activation accounting, router logits, `moe_route`,
18231/// the selection statistics and the `CMF_MOE_TRACE` line — so switching
18232/// executors can never change which experts a token gets.
18233pub(crate) fn moe_ffn_route(
18234    m: &MoeFfn,
18235    x: &[f32],
18236    pool: Option<&Pool>,
18237    allowed: Option<&[bool]>,
18238) -> MoeRoute {
18239    accumulate_act(m, x, 1);
18240    let ne = m.experts.len();
18241    let mut logits = vec![0.0f32; ne];
18242    match &m.resonance {
18243        Some(r) => r.scores(x, &mut logits),
18244        None => m.router.matvec(x, &mut logits, pool),
18245    }
18246    let (idx, p, wsum) = moe_route(&logits, m, allowed);
18247    {
18248        let mut st = m.stats.borrow_mut();
18249        if st.len() < ne {
18250            st.resize(ne, 0);
18251        }
18252        for &e in &idx {
18253            st[e] += 1;
18254        }
18255    }
18256    // `CMF_MOE_TRACE=<file>`: append one line per (layer, token) with the
18257    // selected expert ids. The cumulative `stats` above answer "which
18258    // experts are popular"; a residency design needs the question they
18259    // cannot answer — whether CONSECUTIVE tokens reuse experts (the
18260    // temporal locality an LRU cache lives on, FreeToken §4).
18261    moe_trace(&idx);
18262    MoeRoute {
18263        idx,
18264        p,
18265        wsum,
18266        logits,
18267    }
18268}
18269
18270/// The expert half of `moe_ffn`: run a route's experts on the per-op GPU
18271/// block or the host.
18272pub(crate) fn moe_ffn_experts(
18273    m: &MoeFfn,
18274    x: &[f32],
18275    r: &MoeRoute,
18276    pool: Option<&Pool>,
18277) -> Vec<f32> {
18278    let (idx, p, wsum) = (&r.idx, &r.p, r.wsum);
18279    // D5: the whole layer MoE block in one GPU command buffer (experts — the
18280    // same mmap via a no-copy buffer; intermediate activations on the GPU).
18281    // Same Ffn probe class as the dense chain: one submit per layer
18282    // either wins on this driver stack or it doesn't.
18283    if crate::gpu::enabled_here() {
18284        match crate::gpu::probe_arm(crate::gpu::OpClass::Ffn) {
18285            crate::gpu::ProbeArm::Gpu => {
18286                let t0 = std::time::Instant::now();
18287                if let Some(out) = moe_ffn_gpu(m, x, &idx, &p, wsum, pool) {
18288                    crate::gpu::probe_record(crate::gpu::OpClass::Ffn, true, t0.elapsed());
18289                    return out;
18290                }
18291            }
18292            crate::gpu::ProbeArm::CpuTimed => {
18293                let t0 = std::time::Instant::now();
18294                let out = crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
18295                crate::gpu::probe_record(crate::gpu::OpClass::Ffn, false, t0.elapsed());
18296                return out;
18297            }
18298            crate::gpu::ProbeArm::Cpu => {
18299                return crate::gpu::cpu_scope(|| moe_ffn_cpu(m, x, &idx, &p, wsum, pool));
18300            }
18301        }
18302    }
18303    moe_ffn_cpu(m, x, &idx, &p, wsum, pool)
18304}
18305
18306/// One MoE token through the MiMo expert bank (`crate::mimo_moe`), or —
18307/// when the bank does not serve it — through the host path with the SAME
18308/// route, so the routing statistics and `CMF_MOE_TRACE` see it once.
18309fn moe_ffn_banked(
18310    slot: &mut crate::mimo_moe::Slot,
18311    li: usize,
18312    m: &MoeFfn,
18313    x: &[f32],
18314    pool: Option<&Pool>,
18315) -> Vec<f32> {
18316    let t0 = std::time::Instant::now();
18317    let r = moe_ffn_route(m, x, pool, None);
18318    slot.note_route(t0.elapsed().as_nanos() as u64);
18319    match slot.forward(li, m, x, &r, pool) {
18320        Some(out) => out,
18321        None => crate::qtensor::float_activations_scope(|| {
18322            crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, &r, pool))
18323        }),
18324    }
18325}
18326
18327/// Verify rows share a bank frame; routing and fallback are decode's.
18328fn moe_ffn_banked_rows(
18329    slot: &mut crate::mimo_moe::Slot,
18330    li: usize,
18331    m: &MoeFfn,
18332    xs: &[f32],
18333    b: usize,
18334    hidden: usize,
18335    pool: Option<&Pool>,
18336) -> Vec<f32> {
18337    let t0 = std::time::Instant::now();
18338    let routes: Vec<_> = xs
18339        .chunks_exact(hidden)
18340        .map(|x| moe_ffn_route(m, x, pool, None))
18341        .collect();
18342    slot.note_route(t0.elapsed().as_nanos() as u64);
18343    if let Some(out) = slot.forward_rows(li, m, xs, &routes, pool) {
18344        return out;
18345    }
18346    let mut out = Vec::with_capacity(b * hidden);
18347    for (x, r) in xs.chunks_exact(hidden).zip(&routes) {
18348        let row = slot.forward(li, m, x, r, pool).unwrap_or_else(|| {
18349            // A failed bank must not stream missing experts into the arena.
18350            crate::qtensor::float_activations_scope(|| {
18351                crate::gpu::cpu_scope(|| moe_ffn_experts(m, x, r, pool))
18352            })
18353        });
18354        out.extend(row);
18355    }
18356    out
18357}
18358
18359/// One-shot report of whether the whole-token wgpu graph actually formed.
18360/// A refusal silently reverts to the per-op path, which is how a model can
18361/// look "GPU-accelerated" while every layer walks the host.  A device prefix
18362/// is tracked separately because it still pays a host boundary for the tail.
18363fn graph_note(built: bool, layers_run: usize, total_layers: usize) {
18364    use std::sync::atomic::{AtomicBool, Ordering};
18365    if built {
18366        GRAPH_TOK_OK.fetch_add(1, Ordering::Relaxed);
18367        if total_layers > 0 && layers_run < total_layers {
18368            GRAPH_TOK_PREFIX.fetch_add(1, Ordering::Relaxed);
18369        } else {
18370            GRAPH_TOK_FULL.fetch_add(1, Ordering::Relaxed);
18371        }
18372    } else {
18373        GRAPH_TOK_MISS.fetch_add(1, Ordering::Relaxed);
18374    }
18375    static SAID: AtomicBool = AtomicBool::new(false);
18376    if !SAID.swap(true, Ordering::Relaxed) {
18377        if built {
18378            tracing::info!("wgpu whole-token graph: ACTIVE");
18379        } else {
18380            tracing::warn!("wgpu whole-token graph refused — per-op path");
18381        }
18382    }
18383}
18384
18385/// Whole-token graph outcomes, process-wide: a benchmark that claims a
18386/// GPU number while MISS climbs is measuring the CPU — the honest-bench
18387/// contract makes that an error, not a footnote.
18388pub static GRAPH_TOK_OK: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18389pub static GRAPH_TOK_MISS: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18390/// Graph calls that returned a hidden after running only a leading device
18391/// prefix.  These are valid hybrid executions but must not be reported as a
18392/// full GPU graph in benchmark evidence.
18393pub static GRAPH_TOK_PREFIX: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18394/// Graph calls that covered the complete requested layer span.
18395pub static GRAPH_TOK_FULL: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
18396
18397/// Native Metal TokenGraph completion counters. These are incremented only
18398/// after checked command-buffer completion and successful readback, so a
18399/// fused-head NLL report can prove the route rather than infer it from env.
18400pub static METAL_GRAPH_TOK_OK: std::sync::atomic::AtomicU64 =
18401    std::sync::atomic::AtomicU64::new(0);
18402pub static METAL_GRAPH_HEAD_OK: std::sync::atomic::AtomicU64 =
18403    std::sync::atomic::AtomicU64::new(0);
18404pub static METAL_GRAPH_HEAD_MISS: std::sync::atomic::AtomicU64 =
18405    std::sync::atomic::AtomicU64::new(0);
18406pub static METAL_GRAPH_LAYERS: std::sync::atomic::AtomicU64 =
18407    std::sync::atomic::AtomicU64::new(0);
18408pub static METAL_GRAPH_ERRORS: std::sync::atomic::AtomicU64 =
18409    std::sync::atomic::AtomicU64::new(0);
18410/// Ordinary native-Metal rows-prefill admissions and completed rows.  These
18411/// counters are separate from TokenGraph token/head counts so a batch NLL
18412/// receipt cannot accidentally claim serial execution as batched.
18413pub static METAL_PREFILL_CHUNKS: std::sync::atomic::AtomicU64 =
18414    std::sync::atomic::AtomicU64::new(0);
18415pub static METAL_PREFILL_ROWS: std::sync::atomic::AtomicU64 =
18416    std::sync::atomic::AtomicU64::new(0);
18417pub static METAL_PREFILL_HEAD_ROWS: std::sync::atomic::AtomicU64 =
18418    std::sync::atomic::AtomicU64::new(0);
18419pub static METAL_PREFILL_ERRORS: std::sync::atomic::AtomicU64 =
18420    std::sync::atomic::AtomicU64::new(0);
18421
18422/// `CMF_MOE_BATCH=0` restores the per-expert serial loop — the A/B lever
18423/// for the batched kernel, and how its bit-identity is checked.
18424fn moe_batch_enabled() -> bool {
18425    static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
18426    *ON.get_or_init(|| std::env::var("CMF_MOE_BATCH").as_deref() != Ok("0"))
18427}
18428
18429/// Two-dispatch CPU MoE: every routed expert (and the shared one) fused
18430/// into one gate/up/SiLU dispatch and one down dispatch, instead of two
18431/// pool barriers per expert. Bit-identical to the serial loop below —
18432/// see `moe_gate_up_many` / `moe_down_many`. `None` = the batched kernel
18433/// does not cover this layer, walk the serial path.
18434fn moe_ffn_cpu_batched(
18435    m: &MoeFfn,
18436    x: &[f32],
18437    idx: &[usize],
18438    p: &[f32],
18439    wsum: f32,
18440    pool: Option<&Pool>,
18441) -> Option<Vec<f32>> {
18442    if idx.is_empty() || !moe_batch_enabled() {
18443        return None;
18444    }
18445    // The bake probe reads per-neuron activation mass out of the
18446    // single-expert path; batching would skip it. Rare and offline —
18447    // hand those runs to the serial loop.
18448    if FFN_PROBE.with(|pr| pr.borrow().is_some()) {
18449        return None;
18450    }
18451    let n = idx.len() + usize::from(m.shared.is_some());
18452    let mut pairs = Vec::with_capacity(n);
18453    let mut downs = Vec::with_capacity(n);
18454    let mut ws = Vec::with_capacity(n);
18455    for &e in idx {
18456        let d = &m.experts[e];
18457        if d.act != Act::Silu {
18458            return None;
18459        }
18460        pairs.push((&d.gate_proj, &d.up_proj));
18461        downs.push(&d.down_proj);
18462        ws.push(p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]));
18463    }
18464    // The shared expert goes last, matching the serial loop's order —
18465    // the f32 accumulation order is part of the bit-identity claim.
18466    if let Some((se, gate)) = &m.shared {
18467        if se.act != Act::Silu {
18468            return None;
18469        }
18470        let g = gate.as_ref().map_or(1.0, |gate| {
18471            let mut gl = [0.0f32; 1];
18472            gate.matvec(x, &mut gl, pool);
18473            1.0 / (1.0 + (-gl[0]).exp())
18474        });
18475        pairs.push((&se.gate_proj, &se.up_proj));
18476        downs.push(&se.down_proj);
18477        ws.push(g);
18478    }
18479    let inter = pairs[0].0.rows();
18480    let mut gs: Vec<Vec<f32>> = (0..pairs.len()).map(|_| vec![0f32; inter]).collect();
18481    if !QTensor::moe_gate_up_many(&pairs, x, &mut gs, pool) {
18482        return None;
18483    }
18484    let mut out = attention::take_buf(x.len());
18485    if !QTensor::moe_down_many(&downs, &gs, &ws, &mut out, pool) {
18486        attention::recycle_buf(&mut out);
18487        return None;
18488    }
18489    Some(out)
18490}
18491
18492/// Exact CPU completion for the routed experts a dynamic device cache did
18493/// not contain. The weights are already the router's final normalized mix.
18494/// Keeping this independent of `MoeFfn` makes the job `Sync`: its routing
18495/// statistics live in a `RefCell`, while the immutable expert tensors can be
18496/// evaluated safely in parallel with the GPU's resident subset.
18497pub(crate) fn moe_cold_experts_cpu(
18498    experts: &[(&DenseFfn, f32)],
18499    x: &[f32],
18500    pool: Option<&Pool>,
18501) -> Vec<f32> {
18502    let mut out = attention::take_buf(x.len());
18503    if experts.is_empty() {
18504        return out;
18505    }
18506    let pairs: Vec<_> = experts
18507        .iter()
18508        .map(|(e, _)| (&e.gate_proj, &e.up_proj))
18509        .collect();
18510    let downs: Vec<_> = experts.iter().map(|(e, _)| &e.down_proj).collect();
18511    let weights: Vec<_> = experts.iter().map(|(_, w)| *w).collect();
18512    let inter = experts[0].0.gate_proj.rows();
18513    let mut activations: Vec<Vec<f32>> = (0..experts.len()).map(|_| vec![0.0; inter]).collect();
18514    if QTensor::moe_gate_up_many(&pairs, x, &mut activations, pool)
18515        && QTensor::moe_down_many(&downs, &activations, &weights, &mut out, pool)
18516    {
18517        return out;
18518    }
18519    out.fill(0.0);
18520    for &(expert, weight) in experts {
18521        let mut one = dense_ffn(expert, x, pool);
18522        for (o, v) in out.iter_mut().zip(&one) {
18523            *o += weight * v;
18524        }
18525        attention::recycle_buf(&mut one);
18526    }
18527    out
18528}
18529
18530/// Cold part of a short bank batch. Share each expert's weight stream
18531/// across its tokens, but reduce contributions in each token's route order.
18532/// On an unsupported CPU/layout, retain the single-token cold kernels.
18533pub(crate) fn moe_cold_experts_rows_cpu(
18534    jobs: &[Vec<(&DenseFfn, f32)>],
18535    xs: &[f32],
18536    hidden: usize,
18537    pool: Option<&Pool>,
18538) -> Vec<f32> {
18539    let mut out = vec![0.0; xs.len()];
18540    let mut experts: Vec<&DenseFfn> = Vec::new();
18541    let mut groups: Vec<Vec<usize>> = Vec::new();
18542    let mut terms = vec![Vec::new(); jobs.len()];
18543    for (r, row) in jobs.iter().enumerate() {
18544        for &(e, w) in row {
18545            let g = match experts.iter().position(|&d| std::ptr::eq(d, e)) {
18546                Some(g) => g,
18547                None => {
18548                    experts.push(e);
18549                    groups.push(Vec::new());
18550                    groups.len() - 1
18551                }
18552            };
18553            terms[r].push((g, groups[g].len(), w));
18554            groups[g].push(r);
18555        }
18556    }
18557    if experts.is_empty() {
18558        return out;
18559    }
18560    let pairs: Vec<_> = experts.iter().map(|e| (&e.gate_proj, &e.up_proj)).collect();
18561    let downs: Vec<_> = experts.iter().map(|e| &e.down_proj).collect();
18562    let lens: Vec<_> = groups.iter().map(Vec::len).collect();
18563    let count: usize = lens.iter().sum();
18564    let mut acts = vec![vec![0.0; experts[0].gate_proj.rows()]; count];
18565    let mut ds = vec![vec![0.0; hidden]; count];
18566    if QTensor::moe_gate_up_rows(&pairs, &groups, xs, &mut acts, pool)
18567        && QTensor::moe_down_rows(&downs, &lens, &acts, &mut ds, pool)
18568    {
18569        let mut offset = 0;
18570        let offsets: Vec<_> = lens
18571            .iter()
18572            .map(|&n| {
18573                let start = offset;
18574                offset += n;
18575                start
18576            })
18577            .collect();
18578        for (r, terms) in terms.iter().enumerate() {
18579            for &(g, slot, w) in terms {
18580                for (o, &v) in out[r * hidden..(r + 1) * hidden]
18581                    .iter_mut()
18582                    .zip(&ds[offsets[g] + slot])
18583                {
18584                    *o += w * v;
18585                }
18586            }
18587        }
18588    } else {
18589        for (r, jobs) in jobs.iter().enumerate() {
18590            if !jobs.is_empty() {
18591                let mut row = moe_cold_experts_cpu(jobs, &xs[r * hidden..(r + 1) * hidden], pool);
18592                out[r * hidden..(r + 1) * hidden].copy_from_slice(&row);
18593                attention::recycle_buf(&mut row);
18594            }
18595        }
18596    }
18597    out
18598}
18599
18600/// The pure-CPU MoE expert loop (also the fallback of every GPU refusal).
18601fn moe_ffn_cpu(
18602    m: &MoeFfn,
18603    x: &[f32],
18604    idx: &[usize],
18605    p: &[f32],
18606    wsum: f32,
18607    pool: Option<&Pool>,
18608) -> Vec<f32> {
18609    if let Some(out) = moe_ffn_cpu_batched(m, x, idx, p, wsum, pool) {
18610        return out;
18611    }
18612    let mut out = attention::take_buf(x.len());
18613    for &e in idx {
18614        let mut eo = dense_ffn(&m.experts[e], x, pool);
18615        let w = p[e] / wsum * m.per_expert_scale.as_ref().map_or(1.0, |v| v[e]);
18616        for i in 0..out.len() {
18617            out[i] += w * eo[i];
18618        }
18619        attention::recycle_buf(&mut eo);
18620    }
18621    if let Some((se, gate)) = &m.shared {
18622        let mut so = dense_ffn(se, x, pool);
18623        let g = gate.as_ref().map_or(1.0, |gate| {
18624            let mut gl = [0.0f32; 1];
18625            gate.matvec(x, &mut gl, pool);
18626            1.0 / (1.0 + (-gl[0]).exp())
18627        });
18628        for i in 0..out.len() {
18629            out[i] += g * so[i];
18630        }
18631        attention::recycle_buf(&mut so);
18632    }
18633    out
18634}
18635
18636/// DeepSeek-V2 MLA forward, expand-to-MHA form (see `AttnKind::Mla`):
18637/// per token the latent expands to every head's K/V and the ordinary
18638/// cache + grouped attend do the rest. K head layout is [rope | nope]
18639/// (rotary_dim = qk_rope rotates the shared rope key and each q head's
18640/// prefix); V rows are zero-padded to the K head_dim inside the cache
18641/// and the pad is sliced off before O. Attention importance is not
18642/// accumulated for MLA yet (no eviction interplay).
18643#[allow(clippy::too_many_arguments)]
18644pub(crate) fn mla_attention(
18645    w: &MlaWeights,
18646    normed: &[f32],
18647    cache: &mut crate::kv_cache::LayerKvCache,
18648    position: usize,
18649    inv_freq: &[f32],
18650    rope_scale: f32,
18651    eps: f64,
18652    pool: Option<&Pool>,
18653) -> Vec<f32> {
18654    let (nh, dr, dn, dv, lora) = (w.nh, w.qk_rope, w.qk_nope, w.v_dim, w.lora);
18655    let hd = dr + dn;
18656    let mut q = vec![0.0f32; nh * hd];
18657    match (&w.q_a, &w.q_a_norm) {
18658        (Some(qa), Some(qn)) => {
18659            let mut t = vec![0.0f32; qa.rows()];
18660            qa.matvec(normed, &mut t, pool);
18661            let tn = inference::rms_norm(&t, qn, eps, NormStyle::Qwen);
18662            w.q_proj.matvec(&tn, &mut q, pool);
18663        }
18664        _ => w.q_proj.matvec(normed, &mut q, pool),
18665    }
18666    let mut ca = vec![0.0f32; lora + dr];
18667    w.kv_a.matvec(normed, &mut ca, pool);
18668    let (c_lat, k_rope) = ca.split_at_mut(lora);
18669    let latn = inference::rms_norm(c_lat, &w.kv_a_norm, eps, NormStyle::Qwen);
18670    let mut kvb = vec![0.0f32; nh * (dn + dv)];
18671    w.kv_b.matvec(&latn, &mut kvb, pool);
18672    if !w.nope {
18673        attention::rope_rotate_scaled(k_rope, position, inv_freq, rope_scale);
18674    }
18675    for h in 0..nh {
18676        if !w.nope {
18677            attention::rope_rotate_scaled(
18678                &mut q[h * hd..h * hd + dr],
18679                position,
18680                inv_freq,
18681                rope_scale,
18682            );
18683        }
18684    }
18685    let mut k = vec![0.0f32; nh * hd];
18686    let mut v = vec![0.0f32; nh * hd];
18687    for h in 0..nh {
18688        k[h * hd..h * hd + dr].copy_from_slice(k_rope);
18689        k[h * hd + dr..(h + 1) * hd].copy_from_slice(&kvb[h * (dn + dv)..h * (dn + dv) + dn]);
18690        v[h * hd..h * hd + dv].copy_from_slice(&kvb[h * (dn + dv) + dn..(h + 1) * (dn + dv)]);
18691    }
18692    cache.append(&k, &v, &vec![true; nh]);
18693    let (ao, mut imp) = attention::attend_all_heads(&q, cache, nh, 1, hd, w.scale, None, 0.0);
18694    attention::recycle_buf(&mut imp);
18695    let mut ov = vec![0.0f32; nh * dv];
18696    for h in 0..nh {
18697        ov[h * dv..(h + 1) * dv].copy_from_slice(&ao[h * hd..h * hd + dv]);
18698    }
18699    let mut out = vec![0.0f32; w.o_proj.rows()];
18700    w.o_proj.matvec(&ov, &mut out, pool);
18701    out
18702}
18703
18704/// Gemma-4 dual-branch FFN (spec: see `FfnKind::DenseMoe`). The dense
18705/// branch reads the pre-FFN-normed activation; the router and the
18706/// expert branch read the RAW residual — the router through a
18707/// scale-less rms norm (its constant gain is folded into the weights),
18708/// the experts through `pre_norm_2`. CPU path; GPU graphs refuse the
18709/// layer kind honestly.
18710fn dense_moe_ffn(
18711    dm: &DenseMoeFfn,
18712    x_normed: &[f32],
18713    h_raw: &[f32],
18714    eps: f64,
18715    norm_style: NormStyle,
18716    pool: Option<&Pool>,
18717) -> Vec<f32> {
18718    let mut d = dense_ffn(&dm.dense, x_normed, pool);
18719    d = inference::rms_norm(&d, &dm.post_norm_1, eps, norm_style);
18720    let m = &dm.moe;
18721    let ne = m.experts.len();
18722    let mut logits = vec![0.0f32; ne];
18723    if m.router_input_norm {
18724        let ss: f32 = h_raw.iter().map(|v| v * v).sum::<f32>() / h_raw.len() as f32;
18725        let inv = 1.0 / (ss + eps as f32).sqrt();
18726        let xr: Vec<f32> = h_raw.iter().map(|v| v * inv).collect();
18727        m.router.matvec(&xr, &mut logits, pool);
18728    } else {
18729        m.router.matvec(h_raw, &mut logits, pool);
18730    }
18731    let (idx, p, wsum) = moe_route(&logits, m, None);
18732    {
18733        let mut st = m.stats.borrow_mut();
18734        if st.len() < ne {
18735            st.resize(ne, 0);
18736        }
18737        for &e in &idx {
18738            st[e] += 1;
18739        }
18740    }
18741    let x2 = inference::rms_norm(h_raw, &dm.pre_norm_2, eps, norm_style);
18742    let mo = moe_ffn_cpu(m, &x2, &idx, &p, wsum, pool);
18743    let mo = inference::rms_norm(&mo, &dm.post_norm_2, eps, norm_style);
18744    for (di, mi) in d.iter_mut().zip(&mo) {
18745        *di += mi;
18746    }
18747    d
18748}
18749
18750/// Building the MoE-layer GPU jobs: all selected experts (+shared) must
18751/// be q8_2f-Mapped from the primary mapping; otherwise None → CPU path.
18752/// One-shot report of why the MoE GPU block refused. A silent `?` here
18753/// sends every expert to the CPU with nothing in the logs to say so —
18754/// which is exactly how a q4tp MoE model looked "GPU-accelerated" while
18755/// running entirely on the host.
18756fn moe_gpu_refused(why: &'static str) {
18757    use std::sync::atomic::{AtomicBool, Ordering};
18758    static SAID: AtomicBool = AtomicBool::new(false);
18759    if !SAID.swap(true, Ordering::Relaxed) {
18760        tracing::warn!("MoE GPU block refused ({why}) — experts run on the CPU");
18761    }
18762}
18763
18764fn moe_ffn_gpu(
18765    m: &MoeFfn,
18766    x: &[f32],
18767    idx: &[usize],
18768    p: &[f32],
18769    wsum: f32,
18770    pool: Option<&Pool>,
18771) -> Option<Vec<f32>> {
18772    use crate::gpu::MoeJob;
18773
18774    let mut jobs: Vec<MoeJob> = Vec::with_capacity(idx.len() + 1);
18775    let mut model_ref = None;
18776    for &e in idx {
18777        if moe_push_job(&m.experts[e], x, p[e] / wsum, &mut jobs, &mut model_ref).is_none() {
18778            moe_gpu_refused("push_job(expert)");
18779            return None;
18780        }
18781    }
18782    if let Some((se, gate)) = &m.shared {
18783        let g = gate.as_ref().map_or(1.0, |gate| {
18784            let mut gl = [0.0f32; 1];
18785            gate.matvec(x, &mut gl, pool);
18786            1.0 / (1.0 + (-gl[0]).exp())
18787        });
18788        if moe_push_job(se, x, g, &mut jobs, &mut model_ref).is_none() {
18789            moe_gpu_refused("push_job(shared)");
18790            return None;
18791        }
18792    }
18793    let Some(model) = model_ref else {
18794        moe_gpu_refused("no model_ref");
18795        return None;
18796    };
18797    let hidden = jobs[0].down.1;
18798    let mut out = vec![0.0f32; hidden];
18799    if crate::gpu::moe_block(&model, &jobs, &mut out) {
18800        Some(out)
18801    } else {
18802        moe_gpu_refused("gpu::moe_block");
18803        None
18804    }
18805}
18806
18807/// Single-position FFN dispatch.
18808fn ffn_forward(
18809    ffn: &FfnKind,
18810    x: &[f32],
18811    pool: Option<&Pool>,
18812    experts_allowed: Option<&[bool]>,
18813) -> Vec<f32> {
18814    match ffn {
18815        FfnKind::Dense(d) if !d.segs.is_empty() => tube_ffn(d, x, 1, pool, None),
18816        FfnKind::Dense(d) => dense_ffn(d, x, pool),
18817        FfnKind::Moe(m) => moe_ffn(m, x, pool, experts_allowed),
18818        // Dual-branch layers need the raw residual — their callers
18819        // dispatch dense_moe_ffn directly; the auxiliary paths that land
18820        // here (MTP draft, o1 replay) do not co-occur with gemma-4 MoE.
18821        FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18822    }
18823}
18824
18825/// Fused two-position FFN: gate/up/down streamed once (dense). MoE
18826/// falls back to two singles — expert sets differ per position, there
18827/// is nothing to fuse.
18828fn ffn_forward_pair(
18829    ffn: &FfnKind,
18830    x1: &[f32],
18831    x2: &[f32],
18832    pool: Option<&Pool>,
18833    experts_allowed: Option<&[bool]>,
18834) -> (Vec<f32>, Vec<f32>) {
18835    let d = match ffn {
18836        // A tube layer has nothing to fuse across the pair — the tubes
18837        // are separate matrices; two singles are the honest path.
18838        FfnKind::Dense(d) if !d.segs.is_empty() => {
18839            return (
18840                tube_ffn(d, x1, 1, pool, None),
18841                tube_ffn(d, x2, 1, pool, None),
18842            );
18843        }
18844        FfnKind::Dense(d) => d,
18845        FfnKind::Moe(m) => {
18846            return (
18847                moe_ffn(m, x1, pool, experts_allowed),
18848                moe_ffn(m, x2, pool, experts_allowed),
18849            );
18850        }
18851        FfnKind::DenseMoe(_) => unreachable!("DenseMoe dispatches via dense_moe_ffn"),
18852    };
18853    let inter = d.gate_proj.rows();
18854    FFN_SCRATCH.with(|s| {
18855        let mut s = s.borrow_mut();
18856        let [g1, g2, u1, u2] = &mut *s;
18857        g1.resize(inter, 0.0);
18858        g2.resize(inter, 0.0);
18859        u1.resize(inter, 0.0);
18860        u2.resize(inter, 0.0);
18861        // Multi-matrix pair job: gate+up under one pool dispatch
18862        // (o1s = lane-1 outputs across tensors, o2s = lane-2).
18863        QTensor::matvec2_many(
18864            [&d.gate_proj, &d.up_proj],
18865            x1,
18866            x2,
18867            [g1.as_mut_slice(), u1.as_mut_slice()],
18868            [g2.as_mut_slice(), u2.as_mut_slice()],
18869            pool,
18870        );
18871        for i in 0..inter {
18872            g1[i] = d.act.combine(g1[i], u1[i]);
18873            g2[i] = d.act.combine(g2[i], u2[i]);
18874        }
18875        let mut o1 = attention::take_buf(d.down_proj.rows());
18876        let mut o2 = attention::take_buf(d.down_proj.rows());
18877        d.down_proj.matvec2(g1, g2, &mut o1, &mut o2, pool);
18878        (o1, o2)
18879    })
18880}
18881
18882#[cfg(test)]
18883mod tests {
18884
18885    /// The 0.7.6 prefill-chunk rule: a plain dense stack wholly on a
18886    /// discrete card reads the prompt in wide chunks on x86; every other
18887    /// case keeps the width it had (the GDN-hybrid, MoE and DeepSeek paths
18888    /// were tuned on hardware not measured for this change).
18889    #[test]
18890    fn prefill_chunk_rule_widens_only_dense_on_discrete() {
18891        use super::{
18892            prefill_chunk_rule, ChunkHost, ChunkStackFacts, DISCRETE_DENSE_PREFILL_CHUNK,
18893        };
18894        let dense_card = ChunkStackFacts {
18895            plain_dense: true,
18896            discrete: true,
18897            gpu_on: true,
18898            ..Default::default()
18899        };
18900        assert!(dense_card.dense_on_discrete());
18901        // The bug: a dense Llama on a Vulkan RTX 3090 got 48.
18902        assert_eq!(
18903            prefill_chunk_rule(None, ChunkHost::Other, dense_card.dense_on_discrete()),
18904            DISCRETE_DENSE_PREFILL_CHUNK
18905        );
18906        assert!(DISCRETE_DENSE_PREFILL_CHUNK > 48);
18907        for (label, facts) in [
18908            ("GDN hybrid / MoE / DeepSeek stack", ChunkStackFacts { plain_dense: false, ..dense_card }),
18909            ("integrated GPU", ChunkStackFacts { discrete: false, ..dense_card }),
18910            ("CPU only", ChunkStackFacts { gpu_on: false, discrete: false, ..dense_card }),
18911            ("capacity split", ChunkStackFacts { capacity_split: true, ..dense_card }),
18912            ("multi-GPU plan", ChunkStackFacts { multi_gpu: true, ..dense_card }),
18913            ("O(1) layers", ChunkStackFacts { o1: true, ..dense_card }),
18914        ] {
18915            assert!(!facts.dense_on_discrete(), "{label}");
18916            assert_eq!(
18917                prefill_chunk_rule(None, ChunkHost::Other, facts.dense_on_discrete()),
18918                48,
18919                "{label} keeps the historical x86 chunk"
18920            );
18921        }
18922        // Other hosts are untouched whatever the model.
18923        for dense in [false, true] {
18924            assert_eq!(prefill_chunk_rule(None, ChunkHost::Macos, dense), 512);
18925            assert_eq!(prefill_chunk_rule(None, ChunkHost::Aarch64, dense), 256);
18926        }
18927        // CMF_PREFILL_CHUNK still wins everywhere (and is clamped to ≥ 1).
18928        for host in [ChunkHost::Macos, ChunkHost::Aarch64, ChunkHost::Other] {
18929            for dense in [false, true] {
18930                assert_eq!(prefill_chunk_rule(Some(48), host, dense), 48);
18931                assert_eq!(prefill_chunk_rule(Some(0), host, dense), 1);
18932            }
18933        }
18934    }
18935
18936    #[test]
18937    fn kv_reuse_plan_pulls_rows_decode_wrote_only_on_the_device() {
18938        use super::{ReuseLayer, ReusePlan, kv_reuse_plan};
18939        let full = |host_rows, device_rows| ReuseLayer {
18940            full: true,
18941            host_rows,
18942            device_rows,
18943            device_state: false,
18944        };
18945        // Turn 1: 300-token prompt prefilled on the host, 40 tokens decoded
18946        // by the wgpu graph into the device mirror only. Turn 2 reuses 339.
18947        assert_eq!(
18948            kv_reuse_plan(339, &[full(300, Some(339)), full(300, Some(339))]),
18949            ReusePlan::Pull(vec![(0, 300, 339), (1, 300, 339)])
18950        );
18951        // CPU / Metal: the host owner already holds every forwarded row.
18952        assert_eq!(kv_reuse_plan(339, &[full(339, None)]), ReusePlan::Ready);
18953        // A mirror past the prefix is fine for the host (it gets rewound).
18954        assert_eq!(kv_reuse_plan(339, &[full(339, Some(345))]), ReusePlan::Ready);
18955        // GPU prefix / CPU tail: only the device layers lag.
18956        assert_eq!(
18957            kv_reuse_plan(339, &[full(300, Some(339)), full(339, None)]),
18958            ReusePlan::Pull(vec![(0, 300, 339)])
18959        );
18960        // The device cannot supply the missing rows: never continue.
18961        assert_eq!(kv_reuse_plan(339, &[full(300, Some(320))]), ReusePlan::Fresh);
18962        assert_eq!(kv_reuse_plan(339, &[full(300, None)]), ReusePlan::Fresh);
18963        assert_eq!(kv_reuse_plan(339, &[full(350, None)]), ReusePlan::Fresh);
18964        // A recurrent state advanced on the device cannot be handed to a
18965        // host prefill (it is not rewindable and the host copy is stale).
18966        let conv = |device_state| ReuseLayer {
18967            full: false,
18968            host_rows: 0,
18969            device_rows: None,
18970            device_state,
18971        };
18972        assert_eq!(
18973            kv_reuse_plan(339, &[conv(true), full(300, Some(339))]),
18974            ReusePlan::Fresh
18975        );
18976        assert_eq!(kv_reuse_plan(339, &[conv(false), full(339, None)]), ReusePlan::Ready);
18977    }
18978
18979    #[test]
18980    fn nll_graph_policy_scopes_only_the_fused_head() {
18981        for (label, unmasked, prefer_graph, native_metal, want_graph, want_head) in [
18982            // A Vulkan/Wgpu hidden-only graph remains the quality route.
18983            ("vulkan graph", true, true, false, true, false),
18984            // Native Metal adds the strict fused graph-head contract.
18985            ("native Metal graph", true, true, true, true, true),
18986            // Masked NLL and the explicit non-graph fallback remain unchanged.
18987            ("masked", false, true, false, false, false),
18988            ("graph disabled", true, false, true, false, false),
18989        ] {
18990            let (graph_quality, graph_head_required) =
18991                super::nll_graph_policy(unmasked, prefer_graph, native_metal);
18992            assert_eq!(graph_quality, want_graph, "{label}: graph quality");
18993            assert_eq!(graph_head_required, want_head, "{label}: fused head");
18994        }
18995    }
18996
18997    #[test]
18998    fn mtp_prefill_pair_boundaries_skip_only_final_prompt_row() {
18999        assert_eq!(mtp_prefill_pair_count(0, 128, 256), 128);
19000        assert_eq!(mtp_prefill_pair_count(128, 256, 256), 127);
19001        assert_eq!(mtp_prefill_pair_count(0, 256, 256), 255);
19002        assert_eq!(mtp_prefill_pair_count(256, 256, 256), 0);
19003        assert_eq!(mtp_prefill_pair_count(300, 320, 256), 0);
19004    }
19005
19006    #[test]
19007    fn cancel_flag_stops_generation() {
19008        let mut p = create_test_pipeline(16, 32, 2, 2, 8, 2, 32);
19009        // Set before the call: the prefill loops honour it, the run
19010        // returns immediately with the cancelled reason and no tokens.
19011        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
19012        let r = p.generate_from_ids(&[1, 2, 3], 8, None, None).unwrap();
19013        assert_eq!(r.finish_reason, "cancelled");
19014        assert!(
19015            r.token_ids.is_empty(),
19016            "no tokens after cancel: {:?}",
19017            r.token_ids
19018        );
19019        assert_eq!(p.kv_cache.seq_len(), 0);
19020        assert!(p.kv_history.is_empty());
19021        assert!(!p.graph_want_logits);
19022        assert!(p.graph_logits.is_none());
19023        // Flag auto-cleared: the next call generates normally.
19024        let r2 = p.generate_from_ids(&[1, 2, 3], 4, None, None).unwrap();
19025        assert_ne!(r2.finish_reason, "cancelled");
19026    }
19027    use super::*;
19028
19029    /// sparse_ffn_quant must equal a dense FFN where inactive neurons are
19030    /// zeroed (mask × mmap correctness). On F32 tensors this is EXACT —
19031    /// it validates the row_dot / add_col_scaled / scatter indexing, the
19032    /// bug-prone part. The q8 branches reuse the golden-tested linear
19033    /// The per-token sparse path reads a transposed `down`; it must
19034    /// agree with the arm that computes everything and zeroes the
19035    /// losers, or the speed measurement is measuring a different model.
19036    #[test]
19037    fn dynamic_ffn_equals_the_zeroing_arm() {
19038        let (hidden, inter) = (8usize, 32usize);
19039        let synth = |n: usize, salt: usize| -> Vec<f32> {
19040            (0..n)
19041                .map(|i| (((i * 29 + salt * 13 + 7) % 89) as f32 / 89.0 - 0.5) * 0.6)
19042                .collect()
19043        };
19044        let down = synth(hidden * inter, 3);
19045        let mut down_t = vec![0.0f32; inter * hidden];
19046        for r in 0..hidden {
19047            for c in 0..inter {
19048                down_t[c * hidden + r] = down[r * inter + c];
19049            }
19050        }
19051        let d = DenseFfn {
19052            gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
19053            up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
19054            down_proj: QTensor::from_f32(down.clone(), hidden, inter),
19055            act: Act::Silu,
19056            down_t: Some(QTensor::from_f32(down_t, inter, hidden)),
19057            segs: Vec::new(),
19058        };
19059        let x = synth(hidden, 11);
19060        let k = 12usize;
19061        let got = dense_ffn_dynamic(&d, &x, None, k).expect("down_t present");
19062        // Reference: full compute, keep the k loudest |silu(gate)|.
19063        let mut g = vec![0.0f32; inter];
19064        d.gate_proj.matvec(&x, &mut g, None);
19065        let mut u = vec![0.0f32; inter];
19066        d.up_proj.matvec(&x, &mut u, None);
19067        for v in g.iter_mut() {
19068            *v = inference::silu(*v);
19069        }
19070        keep_top_k(&mut g, k);
19071        for i in 0..inter {
19072            g[i] *= u[i];
19073        }
19074        let mut want = vec![0.0f32; hidden];
19075        d.down_proj.matvec(&g, &mut want, None);
19076        for (a, b) in want.iter().zip(&got) {
19077            assert!((a - b).abs() < 1e-5, "dynamic {b} vs reference {a}");
19078        }
19079    }
19080
19081    /// A tube layer is the same layer, re-cut. With every tube open the
19082    /// answer must equal the dense FFN over the concatenated neurons
19083    /// (the permutation is an identity on the layer's function); with a
19084    /// tube closed it must equal the dense FFN with those neurons
19085    /// zeroed — the mask semantics, now paid for in bytes not read.
19086    #[test]
19087    fn tube_ffn_open_equals_dense_and_closed_equals_masked() {
19088        let (hidden, core, tube) = (8usize, 12usize, 8usize);
19089        let inter = core + tube;
19090        let synth = |n: usize, salt: usize| -> Vec<f32> {
19091            (0..n)
19092                .map(|i| (((i * 41 + salt * 17 + 5) % 97) as f32 / 97.0 - 0.5) * 0.5)
19093                .collect()
19094        };
19095        let (g_all, u_all) = (synth(inter * hidden, 1), synth(inter * hidden, 2));
19096        let d_all = synth(hidden * inter, 3);
19097        // The dense layer, and the same weights cut into core + tube.
19098        let dense = DenseFfn {
19099            gate_proj: QTensor::from_f32(g_all.clone(), inter, hidden),
19100            up_proj: QTensor::from_f32(u_all.clone(), inter, hidden),
19101            down_proj: QTensor::from_f32(d_all.clone(), hidden, inter),
19102            act: Act::Silu,
19103            down_t: None,
19104            segs: Vec::new(),
19105        };
19106        let rows =
19107            |v: &[f32], a: usize, b: usize| -> Vec<f32> { v[a * hidden..b * hidden].to_vec() };
19108        let cols = |v: &[f32], a: usize, b: usize| -> Vec<f32> {
19109            let mut o = Vec::with_capacity(hidden * (b - a));
19110            for r in 0..hidden {
19111                o.extend_from_slice(&v[r * inter + a..r * inter + b]);
19112            }
19113            o
19114        };
19115        let tubed = DenseFfn {
19116            down_t: None,
19117            gate_proj: QTensor::from_f32(rows(&g_all, 0, core), core, hidden),
19118            up_proj: QTensor::from_f32(rows(&u_all, 0, core), core, hidden),
19119            down_proj: QTensor::from_f32(cols(&d_all, 0, core), hidden, core),
19120            act: Act::Silu,
19121            segs: vec![FfnSeg {
19122                gate: QTensor::from_f32(rows(&g_all, core, inter), tube, hidden),
19123                up: QTensor::from_f32(rows(&u_all, core, inter), tube, hidden),
19124                down: QTensor::from_f32(cols(&d_all, core, inter), hidden, tube),
19125                start: core,
19126                width: tube,
19127            }],
19128        };
19129        let x = synth(hidden, 7);
19130        let want = dense_ffn(&dense, &x, None);
19131        let got = tube_ffn(&tubed, &x, 1, None, None);
19132        for (a, b) in want.iter().zip(&got) {
19133            assert!((a - b).abs() < 1e-5, "open tube: {a} vs {b}");
19134        }
19135        // Closed tube: bits on for the core, off for the tube.
19136        let mut bits = vec![0u8; inter.div_ceil(8)];
19137        for n in 0..core {
19138            bits[n / 8] |= 1 << (n % 8);
19139        }
19140        let closed = tube_ffn(&tubed, &x, 1, None, Some(&bits));
19141        let masked = dense_ffn_masked(&dense, &x, None, &bits);
19142        for (a, b) in masked.iter().zip(&closed) {
19143            assert!((a - b).abs() < 1e-5, "closed tube: {a} vs {b}");
19144        }
19145        // The batched arm must agree with the single-position one.
19146        let batch = tube_ffn(&tubed, &x, 1, None, Some(&bits));
19147        for (a, b) in closed.iter().zip(&batch) {
19148            assert_eq!(a, b, "batch arm disagrees with decode arm");
19149        }
19150    }
19151
19152    /// scale, structurally identical to the matvec kernels.
19153    #[test]
19154    fn sparse_ffn_quant_equals_dense_with_inactive_zeroed() {
19155        let (hidden, inter) = (16usize, 40usize);
19156        let synth = |n: usize, salt: usize| -> Vec<f32> {
19157            (0..n)
19158                .map(|i| (((i * 37 + salt * 11 + 3) % 101) as f32 / 101.0 - 0.5) * 0.4)
19159                .collect()
19160        };
19161        let d = DenseFfn {
19162            gate_proj: QTensor::from_f32(synth(inter * hidden, 1), inter, hidden),
19163            up_proj: QTensor::from_f32(synth(inter * hidden, 2), inter, hidden),
19164            down_proj: QTensor::from_f32(synth(hidden * inter, 3), hidden, inter),
19165            act: Act::Silu,
19166            down_t: None,
19167            segs: Vec::new(),
19168        };
19169        let x = synth(hidden, 9);
19170        // Active = every 3rd neuron.
19171        let active: Vec<u16> = (0..inter as u16).filter(|i| i % 3 == 0).collect();
19172
19173        let sparse = sparse_ffn_quant(&d, &x, &active, hidden, None);
19174
19175        // Reference: full dense FFN but g[i]=0 for inactive neurons.
19176        let mut g = vec![0.0f32; inter];
19177        d.gate_proj.matvec(&x, &mut g, None);
19178        let mut u = vec![0.0f32; inter];
19179        d.up_proj.matvec(&x, &mut u, None);
19180        let act_set: std::collections::HashSet<u16> = active.iter().copied().collect();
19181        for i in 0..inter {
19182            g[i] = if act_set.contains(&(i as u16)) {
19183                inference::silu(g[i]) * u[i]
19184            } else {
19185                0.0
19186            };
19187        }
19188        let mut reference = vec![0.0f32; hidden];
19189        d.down_proj.matvec(&g, &mut reference, None);
19190
19191        let max_d = sparse
19192            .iter()
19193            .zip(&reference)
19194            .map(|(a, b)| (a - b).abs())
19195            .fold(0.0f32, f32::max);
19196        assert!(max_d < 1e-5, "sparse != dense-zeroed: max|Δ| = {max_d}");
19197    }
19198
19199    /// Attach a synthetic MTP head (same structure as a main layer).
19200    fn attach_test_mtp(p: &mut Pipeline) {
19201        let (h, inter, heads, kv, hd) = (
19202            p.hidden_size,
19203            p.intermediate_size,
19204            p.num_heads,
19205            p.num_kv_heads,
19206            p.head_dim,
19207        );
19208        let synth = |n: usize, salt: usize| -> Vec<f32> {
19209            (0..n)
19210                .map(|i| (((i * 29 + salt * 23 + 5) % 101) as f32 / 101.0 - 0.5) * 0.2)
19211                .collect()
19212        };
19213        let qt = |rows: usize, cols: usize, salt: usize| -> QTensor {
19214            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19215        };
19216        p.mtp = Some(MtpModule {
19217            enorm: vec![1.0; h],
19218            hnorm: vec![1.0; h],
19219            eh_proj: qt(h, 2 * h, 301),
19220            layer: LayerWeights {
19221                input_norm: vec![1.0; h],
19222                post_norm: vec![1.0; h],
19223                attn_out_norm: None,
19224                ffn_out_norm: None,
19225                layer_scale: None,
19226                ffn: FfnKind::Dense(DenseFfn {
19227                    gate_proj: qt(inter, h, 315),
19228                    up_proj: qt(inter, h, 316),
19229                    down_proj: qt(h, inter, 317),
19230                    act: Act::Silu,
19231                    down_t: None,
19232                    segs: Vec::new(),
19233                }),
19234                attn: AttnKind::Full {
19235                    bias: None,
19236                    wq: qt(heads * hd, h, 311),
19237                    wk: qt(kv * hd, h, 312),
19238                    wv: qt(kv * hd, h, 313),
19239                    wo: qt(h, heads * hd, 314),
19240                    q_norm: None,
19241                    k_norm: None,
19242                    output_gate: false,
19243                    softplus_gate: None,
19244                },
19245            },
19246            final_norm: vec![1.0; h],
19247            kv: crate::kv_cache::LayerKvCache::new(kv, hd),
19248        });
19249    }
19250
19251    #[test]
19252    fn speculative_equals_vanilla_greedy() {
19253        // Speculative decode and the wgpu token graph are mutually
19254        // exclusive; a leaked CMF_GPU=wgpu from a parallel gpu test
19255        // would silently disable drafting. Pin the graph off.
19256        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19257        let run = |spec: bool| {
19258            let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19259            p.sampler_config.temperature = 0.0;
19260            attach_test_mtp(&mut p);
19261            p.speculative = spec;
19262            let r = p.generate("abcdef", 12, None, None).unwrap();
19263            (r.token_ids, r.mtp_drafted, r.mtp_accepted)
19264        };
19265        let (vanilla, d0, _) = run(false);
19266        let (spec, d1, a1) = run(true);
19267        assert_eq!(d0, 0, "vanilla path must not draft");
19268        assert!(d1 > 0, "speculative path must draft");
19269        assert_eq!(
19270            vanilla, spec,
19271            "speculative must reproduce the exact greedy sequence (accepted {a1}/{d1})"
19272        );
19273    }
19274
19275    #[test]
19276    fn speculative_accepts_constant_oracle() {
19277        // See speculative_equals_vanilla_greedy: pin the wgpu graph off.
19278        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
19279        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19280        p.sampler_config.temperature = 0.0;
19281        p.sampler_config.repetition_penalty = 1.0;
19282        // Constant lm_head → every logit equal → both the main model and
19283        // the draft head argmax to token 0: acceptance must be 100%.
19284        p.weights.lm_head = QTensor::from_f32(vec![0.01; 64 * 8], 64, 8);
19285        attach_test_mtp(&mut p);
19286        p.speculative = true;
19287        let r = p.generate("abcd", 10, None, None).unwrap();
19288        assert!(r.mtp_drafted > 0);
19289        assert_eq!(
19290            r.mtp_accepted, r.mtp_drafted,
19291            "constant logits → every draft accepted"
19292        );
19293        // Ties resolve to the same token in both the main and draft
19294        // heads — the sequence is one repeated token.
19295        assert!(r.token_ids.windows(2).all(|w| w[0] == w[1]));
19296    }
19297
19298    #[test]
19299    fn empty_prompt_is_an_error_not_a_panic() {
19300        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
19301        let r = p.generate("", 4, None, None);
19302        assert!(r.is_err(), "empty prompt must be a clean error");
19303    }
19304
19305    #[test]
19306    fn every_token_enters_kv_exactly_once() {
19307        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19308        // Greedy so no RNG variance; byte tokenizer → 3 prompt tokens.
19309        p.sampler_config.temperature = 0.0;
19310        let r = p.generate("abc", 2, None, None).unwrap();
19311        assert_eq!(r.prompt_tokens, 3);
19312        // prompt(3) + first sampled token forwarded before second logits:
19313        // step0 samples from prefill hidden (no extra forward), then
19314        // forwards t1 → cache 4; step1 samples, loop ends (max_tokens).
19315        assert_eq!(
19316            p.kv_cache.seq_len(),
19317            3 + r.tokens_generated - 1,
19318            "each token must be cached exactly once (v1 cached the last prompt token twice)"
19319        );
19320    }
19321
19322    #[test]
19323    fn generation_is_reproducible_with_seed() {
19324        let run = || {
19325            let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19326            p.generate("hello", 8, None, None).unwrap().token_ids
19327        };
19328        assert_eq!(run(), run());
19329    }
19330
19331    #[test]
19332    fn resetting_sampler_restarts_the_seeded_stream() {
19333        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
19334        let config = SamplerConfig {
19335            seed: Some(1234),
19336            ..SamplerConfig::default()
19337        };
19338        p.set_sampler_config(config.clone());
19339        let first = p.generate("hello", 8, None, None).unwrap().token_ids;
19340        p.set_sampler_config(config);
19341        let second = p.generate("hello", 8, None, None).unwrap().token_ids;
19342        assert_eq!(first, second);
19343    }
19344
19345    #[test]
19346    fn eviction_bounds_the_cache() {
19347        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 260);
19348        p.kv_cache.max_seq_len = 6;
19349        p.sampler_config.temperature = 0.0;
19350        let _ = p.generate("abcd", 12, None, None).unwrap();
19351        assert!(
19352            p.kv_cache.seq_len() <= 6 + 1,
19353            "cache must stay bounded by max_seq_len (got {})",
19354            p.kv_cache.seq_len()
19355        );
19356    }
19357
19358    #[test]
19359    fn confidence_matches_tokens_and_is_a_probability() {
19360        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19361        p.sampler_config.temperature = 0.0;
19362        p.sampler_config.repetition_penalty = 1.0;
19363        let r = p.generate("abcd", 10, None, None).unwrap();
19364        assert_eq!(
19365            r.token_confidence.len(),
19366            r.token_ids.len(),
19367            "one confidence per emitted token"
19368        );
19369        for &c in &r.token_confidence {
19370            assert!((0.0..=1.0).contains(&c), "confidence out of [0,1]: {c}");
19371        }
19372        // top1_prob is a valid softmax probability.
19373        let logits = [1.0f32, 3.0, 0.5, 3.0];
19374        let p0 = top1_prob_t(&logits, 1, 1.0);
19375        let p1 = top1_prob_t(&logits, 3, 1.0);
19376        assert!((p0 - p1).abs() < 1e-6, "equal logits → equal prob");
19377        assert!(p0 > 0.0 && p0 < 1.0);
19378        // Calibration temperature > 1 softens an over-confident peak.
19379        let sharp = top1_prob_t(&logits, 1, 1.0);
19380        let soft = top1_prob_t(&logits, 1, 2.0);
19381        assert!(soft < sharp, "higher temperature lowers peak confidence");
19382    }
19383
19384    #[test]
19385    fn trace_is_opt_in_and_parallels_the_output() {
19386        // Off by default: the runtime is silent unless observation asked.
19387        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19388        p.sampler_config.temperature = 0.0;
19389        p.sampler_config.repetition_penalty = 1.0;
19390        let r = p.generate("abcd", 10, None, None).unwrap();
19391        assert!(r.traces.is_empty(), "trace must be empty unless enabled");
19392
19393        // On: exactly one row per emitted token, aligned with the output.
19394        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19395        p.sampler_config.temperature = 0.0;
19396        p.sampler_config.repetition_penalty = 1.0;
19397        p.set_trace(true);
19398        let r = p.generate("abcd", 10, None, None).unwrap();
19399        assert_eq!(r.traces.len(), r.token_ids.len(), "one trace row per token");
19400        for (i, tr) in r.traces.iter().enumerate() {
19401            assert_eq!(tr.t, i, "trace index is sequential");
19402            assert_eq!(tr.token_id, r.token_ids[i], "trace token_id matches output");
19403            assert_eq!(
19404                tr.confidence, r.token_confidence[i],
19405                "trace confidence matches the confidence channel"
19406            );
19407            // No dynamic router in this pipeline → no skill, no coherence.
19408            assert!(tr.active_skill.is_none() && tr.recon.is_none() && !tr.switched);
19409        }
19410    }
19411
19412    #[test]
19413    fn explain_prefill_logits_match_greedy_first_token() {
19414        // `cortiq explain` shows the next-token distribution from
19415        // prefill_next_logits; its argmax must equal what greedy generate
19416        // actually emits first — otherwise explain would lie.
19417        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
19418        p.sampler_config.temperature = 0.0;
19419        p.sampler_config.repetition_penalty = 1.0;
19420        let ids = p.tokenizer.encode("abcd");
19421        let logits = p.prefill_next_logits(&ids, None);
19422        let argmax = logits
19423            .iter()
19424            .enumerate()
19425            .max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
19426            .unwrap()
19427            .0 as u32;
19428        let r = p.generate("abcd", 1, None, None).unwrap();
19429        assert_eq!(
19430            argmax, r.token_ids[0],
19431            "explain preview must match greedy emit"
19432        );
19433    }
19434
19435    #[test]
19436    fn laguna_shared_expert_is_unconditionally_added() {
19437        let matrix = |values: Vec<f32>| QTensor::from_f32(values, 2, 2);
19438        let identity = || matrix(vec![1.0, 0.0, 0.0, 1.0]);
19439        let zero_dense = || DenseFfn {
19440            gate_proj: matrix(vec![0.0; 4]),
19441            up_proj: matrix(vec![0.0; 4]),
19442            down_proj: matrix(vec![0.0; 4]),
19443            act: Act::Silu,
19444            down_t: None,
19445            segs: Vec::new(),
19446        };
19447        let shared = DenseFfn {
19448            gate_proj: identity(),
19449            up_proj: identity(),
19450            down_proj: identity(),
19451            act: Act::Silu,
19452            down_t: None,
19453            segs: Vec::new(),
19454        };
19455        let x = [1.0, 2.0];
19456        let expected = dense_ffn(&shared, &x, None);
19457        let moe = MoeFfn {
19458            router: QTensor::from_f32(vec![0.0, 0.0], 1, 2),
19459            experts: vec![zero_dense()],
19460            top_k: 1,
19461            norm_topk_prob: true,
19462            router_sigmoid: true,
19463            expert_bias: None,
19464            routed_scaling: 1.0,
19465            route_tau: None,
19466            shared: Some((shared, None)),
19467            stats: std::cell::RefCell::new(Vec::new()),
19468            act_sq: std::cell::RefCell::new(Vec::new()),
19469            act_rows: std::cell::RefCell::new(Vec::new()),
19470            mask: None,
19471            per_expert_scale: None,
19472            router_input_norm: false,
19473            resonance: None,
19474            grown: Vec::new(),
19475        };
19476        let actual = moe_ffn_cpu(&moe, &x, &[0], &[0.0], 1.0, None);
19477        for (actual, expected) in actual.iter().zip(expected) {
19478            assert!((actual - expected).abs() < 1e-6);
19479        }
19480    }
19481
19482    /// A tiny MiMo-V2-shaped stack (the M3 fixture): layers [full, sliding,
19483    /// sliding, full]; 4 Q heads over 1 (full) / 2 (sliding) KV heads;
19484    /// head_dim 8 with 4-wide V heads; partial rotary 4 at θ 1e7 (full) /
19485    /// 1e4 (sliding); window 3; learned sinks on the sliding layers; layer
19486    /// 0 a dense FFN, layers 1..3 sigmoid-routed MoE with a selection bias
19487    /// (4 experts, top-2, renormalized, no shared expert). Geometry and
19488    /// sinks go through the same `set_attn_geometry` / `set_layer_sinks`
19489    /// the loader calls.
19490    fn mimo_test_pipeline() -> Pipeline {
19491        let (hs, inter, nh, hd, vd, vocab) = (16usize, 24usize, 4usize, 8usize, 4usize, 64usize);
19492        let kvh = [1usize, 2, 2, 1];
19493        let synth = |n: usize, salt: usize| -> Vec<f32> {
19494            (0..n)
19495                .map(|i| (((i * 31 + salt * 17 + 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
19496                .collect()
19497        };
19498        let qt = |rows: usize, cols: usize, salt: usize| {
19499            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19500        };
19501        let dense = |inter: usize, salt: usize| DenseFfn {
19502            gate_proj: qt(inter, hs, salt),
19503            up_proj: qt(inter, hs, salt + 1),
19504            down_proj: qt(hs, inter, salt + 2),
19505            act: Act::Silu,
19506            down_t: None,
19507            segs: Vec::new(),
19508        };
19509        let layers: Vec<LayerWeights> = (0..4)
19510            .map(|li| LayerWeights {
19511                input_norm: vec![1.0; hs],
19512                post_norm: vec![1.0; hs],
19513                attn_out_norm: None,
19514                ffn_out_norm: None,
19515                layer_scale: None,
19516                ffn: if li == 0 {
19517                    FfnKind::Dense(dense(inter, 50))
19518                } else {
19519                    FfnKind::Moe(MoeFfn {
19520                        router: qt(4, hs, 60 + li),
19521                        experts: (0..4).map(|e| dense(8, 70 + li * 10 + e * 3)).collect(),
19522                        top_k: 2,
19523                        norm_topk_prob: true,
19524                        router_sigmoid: true,
19525                        expert_bias: Some(vec![0.02, -0.03, 0.01, 0.0]),
19526                        routed_scaling: 1.0,
19527                        route_tau: None,
19528                        shared: None,
19529                        stats: std::cell::RefCell::new(Vec::new()),
19530                        act_sq: std::cell::RefCell::new(Vec::new()),
19531                        act_rows: std::cell::RefCell::new(Vec::new()),
19532                        mask: None,
19533                        per_expert_scale: None,
19534                        router_input_norm: false,
19535                        resonance: None,
19536                        grown: Vec::new(),
19537                    })
19538                },
19539                attn: AttnKind::Full {
19540                    wq: qt(nh * hd, hs, li * 10 + 1),
19541                    wk: qt(kvh[li] * hd, hs, li * 10 + 2),
19542                    wv: qt(kvh[li] * vd, hs, li * 10 + 3),
19543                    wo: qt(hs, nh * vd, li * 10 + 4),
19544                    q_norm: None,
19545                    k_norm: None,
19546                    output_gate: false,
19547                    softplus_gate: None,
19548                    bias: None,
19549                },
19550            })
19551            .collect();
19552        let mut p = Pipeline::new(
19553            Tokenizer::byte_level(),
19554            PipelineWeights {
19555                embed_tokens: qt(vocab, hs, 100),
19556                layers,
19557                lm_head: qt(vocab, hs, 200),
19558                final_norm: vec![1.0; hs],
19559            },
19560            hs,
19561            inter,
19562            nh,
19563            1, // header num_kv_heads (the full layers')
19564            hd,
19565            4,
19566            4,
19567            false,
19568            vocab,
19569            1e-6,
19570            1e7,
19571            NormStyle::Qwen,
19572            4096,
19573            SamplerConfig {
19574                seed: Some(7),
19575                ..Default::default()
19576            },
19577        );
19578        // Diagnostics stay off whatever the test environment exports.
19579        p.layer_dump = None;
19580        p.set_rotary(4, 1e7);
19581        p.sliding_layers = Some(vec![false, true, true, false]);
19582        p.swa = Some((3, usize::MAX));
19583        p.rotary_dim_local = Some(4);
19584        p.inv_freq_local = Some(std::sync::Arc::new(attention::rope_inv_freq(4, 1e4)));
19585        p.set_attn_geometry(Some(kvh.to_vec()), Some(vd)).unwrap();
19586        p.set_layer_sinks(1, vec![0.5, -1.0, 1.5, 0.0]).unwrap();
19587        p.set_layer_sinks(2, vec![-0.25, 2.0, 0.75, -1.5]).unwrap();
19588        p
19589    }
19590
19591    #[test]
19592    fn mimo_embedded_prompt_uses_rows_and_never_reuses_token_only_kv() {
19593        let mut p = mimo_test_pipeline();
19594        p.speculative = false;
19595        p.ignore_eos = true;
19596        p.sampler_config.temperature = 0.0;
19597        p.sampler_config.repetition_penalty = 1.0;
19598        let a = vec![3, 5, 7, 9, 11, 13];
19599        let b = vec![4, 8, 12, 16, 20, 24];
19600        let rows: Vec<_> = b.iter().flat_map(|&id| p.embed_id(id)).collect();
19601        let expected = p.generate_from_ids(&b, 8, None, None).unwrap().token_ids;
19602        // Same placeholder IDs as an earlier request are not a cache key
19603        // for different media. The actual rows, not a re-embedding of a,
19604        // must determine the continuation.
19605        let actual = p.generate_from_embeds(&a, &rows, 8, None, None).unwrap().token_ids;
19606        assert_eq!(actual, expected);
19607        assert!(p.kv_history.is_empty());
19608        let mut extended = a.clone();
19609        extended.push(17);
19610        let after_media = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19611        p.reset_session();
19612        let fresh = p.generate_from_ids(&extended, 8, None, None).unwrap().token_ids;
19613        assert_eq!(after_media, fresh);
19614        assert!(p.generate_from_embeds(&a, &rows[..rows.len()-1], 1, None, None).is_err());
19615        // Force a real token-prefix reuse opportunity into the media call.
19616        // Those labels are unchanged, but their embeddings now describe a
19617        // different source sequence and every KV row must be rebuilt.
19618        p.reset_session();
19619        p.generate_from_ids(&a, 1, None, None).unwrap();
19620        let mut media_ids = p.kv_history.clone();
19621        assert!(!media_ids.is_empty());
19622        media_ids.extend_from_slice(&[19, 21, 23]);
19623        let source_ids: Vec<_> = (0..media_ids.len()).map(|i| b[i % b.len()]).collect();
19624        let source_rows: Vec<_> = source_ids.iter().flat_map(|&id| p.embed_id(id)).collect();
19625        let mut oracle = mimo_test_pipeline();
19626        oracle.speculative = false;
19627        oracle.ignore_eos = true;
19628        oracle.sampler_config.temperature = 0.0;
19629        oracle.sampler_config.repetition_penalty = 1.0;
19630        let expected = oracle.generate_from_ids(&source_ids, 8, None, None).unwrap().token_ids;
19631        assert_eq!(p.generate_from_embeds(&media_ids, &source_rows, 8, None, None).unwrap().token_ids, expected);
19632        assert!(p.kv_history.is_empty());
19633        let mut bad = rows;
19634        bad[0] = f32::NAN;
19635        assert!(p.generate_from_embeds(&a, &bad, 1, None, None).is_err());
19636    }
19637
19638    fn f32_bits(v: &[f32]) -> Vec<u32> {
19639        v.iter().map(|x| x.to_bits()).collect()
19640    }
19641
19642    /// M3 acceptance: on the MiMo-shaped stack the decode walk (one
19643    /// position at a time through `forward_layers`) and the batched
19644    /// prefill (`prefill_batch_span`, whole prompt and split in two
19645    /// chunks) give bit-identical logits at all 12 positions — per-layer
19646    /// KV heads, narrow V, sinks, the window and the biased sigmoid MoE all
19647    /// agree across the two walks.
19648    #[test]
19649    fn mimo_shaped_decode_matches_prefill_batch_bitwise() {
19650        let mut p = mimo_test_pipeline();
19651        let kv: Vec<usize> = p.kv_cache.layers.iter().map(|l| l.num_kv_heads).collect();
19652        assert_eq!(kv, vec![1, 2, 2, 1]);
19653        assert!(p.kv_cache.layers[0].sinks.is_none() && p.kv_cache.layers[3].sinks.is_none());
19654        assert!(p.kv_cache.layers[1].sinks.is_some() && p.kv_cache.layers[2].sinks.is_some());
19655        let ids: Vec<u32> = (0..12u32).map(|i| (i * 7 + 3) % 64).collect();
19656        let hs = p.hidden_size;
19657        let mut decode = Vec::new();
19658        for (pos, &id) in ids.iter().enumerate() {
19659            let e = p.embed_single(id);
19660            let h = p.forward_layers(&e, pos, None);
19661            decode.push(p.logits_from_hidden(&h));
19662        }
19663        for l in &p.kv_cache.layers {
19664            assert_eq!(l.seq_len, 12);
19665            // V rows are padded to head_dim inside the cache.
19666            assert_eq!(l.head_values(0).len(), 12 * 8);
19667        }
19668        assert!(decode.iter().flatten().all(|v| v.is_finite()));
19669
19670        p.clear_sequence_state();
19671        let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19672        for pos in 0..ids.len() {
19673            let lg = p.logits_from_hidden(&hb[pos * hs..(pos + 1) * hs]);
19674            assert_eq!(
19675                f32_bits(&decode[pos]),
19676                f32_bits(&lg),
19677                "whole prompt, pos {pos}"
19678            );
19679        }
19680
19681        p.clear_sequence_state();
19682        let a = p.prefill_batch_span(PrefillIn::Ids(&ids[..5]), 0, None, 0, p.num_layers);
19683        let b = p.prefill_batch_span(PrefillIn::Ids(&ids[5..]), 5, None, 0, p.num_layers);
19684        for pos in 0..ids.len() {
19685            let row = if pos < 5 {
19686                &a[pos * hs..(pos + 1) * hs]
19687            } else {
19688                &b[(pos - 5) * hs..(pos - 4) * hs]
19689            };
19690            let lg = p.logits_from_hidden(row);
19691            assert_eq!(
19692                f32_bits(&decode[pos]),
19693                f32_bits(&lg),
19694                "two chunks, pos {pos}"
19695            );
19696        }
19697
19698        // The fixture is not degenerate: the sinks and the window each
19699        // change the answer.
19700        let last = |p: &mut Pipeline| {
19701            p.clear_sequence_state();
19702            let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19703            p.logits_from_hidden(&hb[11 * hs..12 * hs])
19704        };
19705        let base = last(&mut p);
19706        let mut no_sinks = mimo_test_pipeline();
19707        for l in &mut no_sinks.kv_cache.layers {
19708            l.sinks = None;
19709        }
19710        assert_ne!(
19711            f32_bits(&last(&mut no_sinks)),
19712            f32_bits(&base),
19713            "sinks are live"
19714        );
19715        let mut wide = mimo_test_pipeline();
19716        wide.swa = Some((64, usize::MAX));
19717        assert_ne!(
19718            f32_bits(&last(&mut wide)),
19719            f32_bits(&base),
19720            "window is live"
19721        );
19722
19723        // Generation runs end to end on the same stack.
19724        p.clear_sequence_state();
19725        p.ignore_eos = true;
19726        let r = p.generate_from_ids(&ids, 4, None, None).unwrap();
19727        assert_eq!(r.token_ids.len(), 4);
19728    }
19729
19730    /// A Spark-X2.5-shaped stack: layers 0..3 slide (window 6), layer 3
19731    /// is global, every attention carries the head-wise sigmoid g_proj gate.
19732    fn spark_test_pipeline(trim: Option<(usize, usize)>) -> Pipeline {
19733        let (hs, nh) = (16usize, 4usize);
19734        let mut p = create_test_pipeline(hs, 24, nh, 2, 8, 4, 64);
19735        p.layer_dump = None;
19736        p.swa = Some((6, 4));
19737        p.proj_gate_sigmoid = true;
19738        for (li, lw) in p.weights.layers.iter_mut().enumerate() {
19739            if let AttnKind::Full { softplus_gate, .. } = &mut lw.attn {
19740                let g: Vec<f32> = (0..nh * hs)
19741                    .map(|i| (((i * 29 + li * 7) % 83) as f32 / 83.0 - 0.5) * 0.8)
19742                    .collect();
19743                *softplus_gate = Some((QTensor::from_f32(g, nh, hs), true));
19744            }
19745        }
19746        p.swa_trim = trim;
19747        p
19748    }
19749
19750    /// Trimming the sliding tails changes nothing but memory: the same
19751    /// prompt in uneven chunks, a decode walk, and fused pairs with a
19752    /// 2-row rollback give bit-identical hiddens with and without it,
19753    /// while the trimmed sliding layers stay bounded and the global layer
19754    /// keeps every row.
19755    #[test]
19756    fn swa_trim_pipeline_matches_untrimmed_bitwise() {
19757        let mut a = spark_test_pipeline(None);
19758        let mut b = spark_test_pipeline(Some((2, 4)));
19759        assert_eq!(
19760            (0..4).map(|li| b.layer_window(li)).collect::<Vec<_>>(),
19761            vec![Some(6), Some(6), Some(6), None]
19762        );
19763        let ids: Vec<u32> = (0..41u32).map(|i| (i * 11 + 5) % 64).collect();
19764        let mut pos = 0usize;
19765        for &n in [5usize, 13, 7, 16].iter().cycle() {
19766            if pos >= ids.len() {
19767                break;
19768            }
19769            let end = (pos + n).min(ids.len());
19770            let ha = a.prefill_batch_span(PrefillIn::Ids(&ids[pos..end]), pos, None, 0, 4);
19771            let hb = b.prefill_batch_span(PrefillIn::Ids(&ids[pos..end]), pos, None, 0, 4);
19772            assert_eq!(f32_bits(&ha), f32_bits(&hb), "prefill chunk at {pos}");
19773            pos = end;
19774        }
19775        for _ in 0..23 {
19776            let e = a.embed_single(((pos * 7) % 64) as u32);
19777            let ha = a.forward_layers(&e, pos, None);
19778            let hb = b.forward_layers(&e, pos, None);
19779            assert_eq!(f32_bits(&ha), f32_bits(&hb), "decode at {pos}");
19780            for li in 0..3 {
19781                assert!(b.kv_cache.layers[li].seq_len <= 12, "layer {li} bounded");
19782            }
19783            pos += 1;
19784        }
19785        for round in 0..9 {
19786            let (e1, e2) = (a.embed_single(round * 3 + 1), a.embed_single(round * 5 + 2));
19787            let (a1, a2) = a.forward_pair(&e1, &e2, pos);
19788            let (b1, b2) = b.forward_pair(&e1, &e2, pos);
19789            assert_eq!(f32_bits(&a1), f32_bits(&b1), "pair lane 1 round {round}");
19790            assert_eq!(f32_bits(&a2), f32_bits(&b2), "pair lane 2 round {round}");
19791            if round % 2 == 1 {
19792                // A rejected draft: both lanes roll back.
19793                for p in [&mut a, &mut b] {
19794                    for l in &mut p.kv_cache.layers {
19795                        l.truncate_last(2);
19796                    }
19797                }
19798            } else {
19799                pos += 2;
19800            }
19801        }
19802        let e = a.embed_single(9);
19803        let (ha, hb) = (
19804            a.forward_layers(&e, pos, None),
19805            b.forward_layers(&e, pos, None),
19806        );
19807        assert_eq!(f32_bits(&ha), f32_bits(&hb), "after the pairs");
19808        for li in 0..3 {
19809            let (la, lb) = (&a.kv_cache.layers[li], &b.kv_cache.layers[li]);
19810            assert!(lb.base() > 0, "layer {li} trimmed");
19811            assert_eq!(la.base(), 0);
19812            assert_eq!(lb.pos_len(), la.seq_len);
19813            assert_eq!(lb.head_keys(0), &la.head_keys(0)[lb.base() * 8..]);
19814        }
19815        let (ga, gb) = (&a.kv_cache.layers[3], &b.kv_cache.layers[3]);
19816        assert_eq!(
19817            (gb.base(), gb.seq_len),
19818            (0, ga.seq_len),
19819            "the global layer keeps all"
19820        );
19821        assert_eq!(b.kv_cache.seq_len(), a.kv_cache.seq_len());
19822        assert!(b.kv_cache.total_memory_bytes() < a.kv_cache.total_memory_bytes());
19823    }
19824
19825    /// The network split's KvFetch hand-off on a trimmed stack: every
19826    /// layer shipped over the wire mid-sequence (the sliding ones as
19827    /// `FullTail`) into a fresh pipeline that never ran a position, which
19828    /// then decodes on — bit for bit what the untrimmed pipeline decodes.
19829    #[test]
19830    fn swa_trim_wire_handoff_continues_bitwise() {
19831        let mut a = spark_test_pipeline(None);
19832        let mut b = spark_test_pipeline(Some((2, 4)));
19833        let ids: Vec<u32> = (0..29u32).map(|i| (i * 13 + 7) % 64).collect();
19834        let ha = a.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, 4);
19835        let hb = b.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, 4);
19836        assert_eq!(f32_bits(&ha), f32_bits(&hb));
19837        let mut pos = ids.len();
19838        for _ in 0..5 {
19839            let e = a.embed_single(((pos * 5) % 64) as u32);
19840            assert_eq!(
19841                f32_bits(&a.forward_layers(&e, pos, None)),
19842                f32_bits(&b.forward_layers(&e, pos, None))
19843            );
19844            pos += 1;
19845        }
19846        let mut c = spark_test_pipeline(Some((2, 4)));
19847        for li in 0..4 {
19848            let bytes = b.kv_cache.layers[li].export_wire(false).unwrap();
19849            c.kv_cache.layers[li].import_wire(&bytes).unwrap();
19850        }
19851        assert!(c.kv_cache.layers[0].base() > 0, "a tail travelled");
19852        assert_eq!(c.kv_cache.seq_len(), pos);
19853        for step in 0..17 {
19854            let e = a.embed_single(((pos * 3 + 1) % 64) as u32);
19855            let ha = a.forward_layers(&e, pos, None);
19856            let hc = c.forward_layers(&e, pos, None);
19857            assert_eq!(
19858                f32_bits(&ha),
19859                f32_bits(&hc),
19860                "after the hand-off, step {step}"
19861            );
19862            pos += 1;
19863        }
19864        assert!(c.kv_cache.layers[0].seq_len <= 12, "the receiver trims on");
19865    }
19866
19867    /// A synthetic MiMo draft stack of `n` layers for `mimo_test_pipeline`
19868    /// (the SWA geometry of its sliding layers: 2 KV heads, head 8 / V 4).
19869    fn mimo_test_mtp(n: usize, gain: f32) -> mimo_mtp::MimoMtp {
19870        let (hs, inter, nh, hd, vd, nkv) = (16usize, 24usize, 4usize, 8usize, 4usize, 2usize);
19871        let synth = |len: usize, salt: usize| -> Vec<f32> {
19872            (0..len)
19873                .map(|i| (((i * 37 + salt * 13 + 3) % 89) as f32 / 89.0 - 0.5) * 0.6 * gain)
19874                .collect()
19875        };
19876        let qt = |rows: usize, cols: usize, salt: usize| {
19877            QTensor::from_f32(synth(rows * cols, salt), rows, cols)
19878        };
19879        let layers = (0..n)
19880            .map(|k| {
19881                let s = 500 + k * 40;
19882                let mut kv = crate::kv_cache::LayerKvCache::new(nkv, hd);
19883                kv.sinks = Some(vec![0.3, -0.7, 1.1, 0.0]);
19884                MtpModule {
19885                    enorm: vec![1.0; hs],
19886                    hnorm: vec![1.0; hs],
19887                    eh_proj: qt(hs, 2 * hs, s),
19888                    layer: LayerWeights {
19889                        input_norm: vec![1.0; hs],
19890                        post_norm: vec![1.0; hs],
19891                        attn_out_norm: None,
19892                        ffn_out_norm: None,
19893                        layer_scale: None,
19894                        attn: AttnKind::Full {
19895                            wq: qt(nh * hd, hs, s + 1),
19896                            wk: qt(nkv * hd, hs, s + 2),
19897                            wv: qt(nkv * vd, hs, s + 3),
19898                            wo: qt(hs, nh * vd, s + 4),
19899                            q_norm: None,
19900                            k_norm: None,
19901                            output_gate: false,
19902                            softplus_gate: None,
19903                            bias: None,
19904                        },
19905                        ffn: FfnKind::Dense(DenseFfn {
19906                            gate_proj: qt(inter, hs, s + 5),
19907                            up_proj: qt(inter, hs, s + 6),
19908                            down_proj: qt(hs, inter, s + 7),
19909                            act: Act::Silu,
19910                            down_t: None,
19911                            segs: Vec::new(),
19912                        }),
19913                    },
19914                    final_norm: vec![1.0; hs],
19915                    kv,
19916                }
19917            })
19918            .collect();
19919        mimo_mtp::MimoMtp::from_layers(layers)
19920    }
19921
19922    fn mimo_greedy(p: &mut Pipeline, ids: &[u32], n: usize, spec: bool) -> GenerateResult {
19923        p.clear_sequence_state();
19924        p.speculative = spec;
19925        p.ignore_eos = true;
19926        p.sampler_config.temperature = 0.0;
19927        p.generate_from_ids(ids, n, None, None).unwrap()
19928    }
19929
19930    /// The draft stack's incremental rounds (a few rows per layer, last
19931    /// round's provisional rows dropped) give exactly the teacher-forced
19932    /// table of one causal pass per layer over the whole sequence — the
19933    /// table `tools/mimo_ref.py mtp` computes for variant A: layer k, row
19934    /// j reads (x[j+k+1], norm(h_j)) at RoPE position j.
19935    #[test]
19936    fn mimo_mtp_incremental_rounds_equal_the_teacher_forced_table() {
19937        // Both readings of the backbone hidden: pre-final-norm (default)
19938        // and post-final-norm (`CMF_MIMO_MTP_HIDDEN=post`).
19939        for post in [false, true] {
19940            let mut p = mimo_test_pipeline();
19941            // A non-trivial final norm, so the two readings differ.
19942            p.weights.final_norm = (0..p.hidden_size).map(|i| 0.5 + 0.1 * i as f32).collect();
19943            let mut st0 = mimo_test_mtp(3, 1.0);
19944            st0.post_norm_hidden = post;
19945            p.mimo_mtp = Some(st0);
19946            let ids: Vec<u32> = (0..14u32).map(|i| (i * 11 + 5) % 64).collect();
19947            let hs = p.hidden_size;
19948            let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
19949            p.mimo_note_rows(&hb, 0);
19950            let mut st = p.mimo_mtp.take().unwrap();
19951            // Incremental: one round per t through the decode path (later
19952            // tokens from `ids`, the probe's teacher forcing).
19953            let k = 3;
19954            let mut inc = Vec::new();
19955            for t in 0..ids.len() - k - 1 {
19956                inc.push(p.mimo_mtp_draft(&mut st, t, &ids, k));
19957            }
19958            // Reference: per layer, ONE batched causal pass over all rows
19959            // with fresh caches.
19960            let s = ids.len();
19961            let mut reference = vec![vec![0u32; k]; s - k - 1];
19962            let mut fresh = mimo_test_mtp(3, 1.0);
19963            for (layer, m) in fresh.layers.iter_mut().enumerate() {
19964                let n = s - layer - 1;
19965                let mut cats = vec![0.0f32; n * 2 * hs];
19966                for j in 0..n {
19967                    let e = p.embed_single(ids[j + layer + 1]);
19968                    let raw = &hb[j * hs..(j + 1) * hs];
19969                    let g = if post {
19970                        inference::rms_norm(raw, &p.weights.final_norm, p.rms_eps, p.norm_style)
19971                    } else {
19972                        raw.to_vec()
19973                    };
19974                    let (ce, ch) = cats[j * 2 * hs..(j + 1) * 2 * hs].split_at_mut(hs);
19975                    inference::rms_norm_into(&e, &m.enorm, p.rms_eps, p.norm_style, ce);
19976                    inference::rms_norm_into(&g, &m.hnorm, p.rms_eps, p.norm_style, ch);
19977                }
19978                let mut x = vec![0.0f32; n * hs];
19979                m.eh_proj.matmat(&cats, n, &mut x, None);
19980                p.mimo_mtp_block(m, &mut x, n, 0);
19981                for (t, row) in reference.iter_mut().enumerate() {
19982                    let y = inference::rms_norm(
19983                        &x[t * hs..(t + 1) * hs],
19984                        &m.final_norm,
19985                        p.rms_eps,
19986                        p.norm_style,
19987                    );
19988                    row[layer] = sampler::argmax(&p.lm_head_forward(&y));
19989                }
19990            }
19991            assert_eq!(inc, reference, "post_norm_hidden = {post}");
19992            // Not a degenerate table: the drafts vary.
19993            let distinct: std::collections::HashSet<u32> =
19994                inc.iter().flatten().copied().collect();
19995            assert!(distinct.len() > 3, "{inc:?}");
19996            // Each layer's cache ends holding rows up to the last round start.
19997            let last_t = ids.len() - k - 2;
19998            for m in &st.layers {
19999                assert_eq!(m.kv.seq_len, last_t + 1);
20000            }
20001        }
20002    }
20003
20004    /// Greedy with the MiMo draft stack is the plain greedy stream, token
20005    /// for token — with the real draft layers (low acceptance) and with a
20006    /// drafter that is right most of the time (exercises accepted prefixes
20007    /// of every length, the KV truncation of the rejected rows and the
20008    /// logits hand-off to the loop top), under the default repetition
20009    /// penalty.
20010    #[test]
20011    fn mimo_speculative_greedy_equals_plain_greedy() {
20012        unsafe { std::env::set_var("CMF_GPU_WGPU_GRAPH", "0") };
20013        let ids: Vec<u32> = (0..9u32).map(|i| (i * 7 + 3) % 64).collect();
20014        let n = 24;
20015        let mut p = mimo_test_pipeline();
20016        let plain = mimo_greedy(&mut p, &ids, n, false);
20017        assert_eq!(plain.mtp_drafted, 0);
20018        assert_eq!(plain.token_ids.len(), n);
20019        let plain_kv = p.kv_cache.layers[0].seq_len;
20020
20021        // Real draft layers.
20022        p.mimo_mtp = Some(mimo_test_mtp(3, 1.0));
20023        let spec = mimo_greedy(&mut p, &ids, n, true);
20024        assert!(spec.mtp_drafted > 0, "the round must draft");
20025        assert_eq!(spec.token_ids, plain.token_ids);
20026        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
20027
20028        // A drafter reading the true continuation with every fifth token
20029        // wrong: accepted prefixes of 0..=3 all occur.
20030        let mut truth: Vec<u32> = ids.clone();
20031        truth.extend(&plain.token_ids);
20032        let mut noisy = truth.clone();
20033        for (i, t) in noisy.iter_mut().enumerate() {
20034            if i % 5 == 0 {
20035                *t = (*t + 1) % 64;
20036            }
20037        }
20038        let mut st = mimo_test_mtp(3, 1.0);
20039        st.draft_override = Some(noisy);
20040        p.mimo_mtp = Some(st);
20041        let spec = mimo_greedy(&mut p, &ids, n, true);
20042        assert_eq!(spec.token_ids, plain.token_ids);
20043        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
20044        let stats = p.mimo_mtp.as_ref().unwrap().stats.clone();
20045        assert_eq!(stats.accepted as usize, spec.mtp_accepted);
20046        assert!(spec.mtp_accepted > 0 && spec.mtp_accepted < spec.mtp_drafted);
20047        assert!(
20048            stats.accept_hist.iter().filter(|&&c| c > 0).count() >= 3,
20049            "{:?}",
20050            stats.accept_hist
20051        );
20052        assert!(stats.tokens_per_round() > 1.5, "{}", stats.line());
20053
20054        // A perfect drafter: every draft accepted, rounds of K+1 tokens,
20055        // and the budget is never overrun.
20056        let mut st = mimo_test_mtp(3, 1.0);
20057        st.draft_override = Some(truth);
20058        p.mimo_mtp = Some(st);
20059        let spec = mimo_greedy(&mut p, &ids, n, true);
20060        assert_eq!(spec.token_ids, plain.token_ids);
20061        assert_eq!(spec.mtp_accepted, spec.mtp_drafted);
20062        assert_eq!(p.kv_cache.layers[0].seq_len, plain_kv);
20063
20064        // CMF_MTP=0 path: the stack is attached but idle.
20065        let off = mimo_greedy(&mut p, &ids, n, false);
20066        assert_eq!(off.token_ids, plain.token_ids);
20067        assert_eq!(off.mtp_drafted, 0);
20068    }
20069
20070    /// The wgpu graphs carry MiMo-V2's attention per layer (KV heads,
20071    /// narrow V, sinks, windows, two RoPE tables): no attention-level
20072    /// decline for it any more, and the geometry each layer hands the
20073    /// graph is exactly what the CPU attention reads for that layer. The
20074    /// descriptive reasons stay (the Metal graphs and the q1 dropin still
20075    /// decline on them), and what the per-layer geometry cannot express
20076    /// keeps a named wgpu decline.
20077
20078    #[test]
20079    fn mimo_shaped_model_rides_the_wgpu_graph_geometry() {
20080        let p = mimo_test_pipeline();
20081        assert_eq!(
20082            p.graph_attn_decline_reason(),
20083            Some("per-layer KV head counts")
20084        );
20085        assert_eq!(p.wgpu_graph_attn_decline(), None);
20086        let g0 = p.graph_attn_geom(0).expect("full layer geometry");
20087        assert_eq!(
20088            (g0.nkv, g0.dv, g0.rd, g0.window, g0.sink.is_some()),
20089            (1, 4, 4, None, false)
20090        );
20091        assert_eq!(g0.invf, p.inv_freq.as_slice());
20092        let g1 = p.graph_attn_geom(1).expect("sliding layer geometry");
20093        assert_eq!((g1.nkv, g1.dv, g1.rd, g1.window), (2, 4, 4, Some(3)));
20094        assert_eq!(g1.sink, Some(&[0.5f32, -1.0, 1.5, 0.0][..]));
20095        assert_eq!(g1.invf, p.inv_freq_local.as_ref().unwrap().as_slice());
20096        assert_ne!(g0.invf, g1.invf, "two RoPE tables");
20097        let g3 = p.graph_attn_geom(3).expect("full layer geometry");
20098        assert_eq!((g3.nkv, g3.window, g3.sink.is_some()), (1, None, false));
20099
20100        // No wgpu device in this process: the builders run and decline on
20101        // the (f32, unmapped) experts — never with an attention line.
20102        let emb = p.embed_single(3);
20103        let mut lg = Vec::new();
20104        assert!(
20105            p.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, p.num_layers)
20106                .is_none()
20107        );
20108        let mut hid = emb.clone();
20109        assert_eq!(
20110            p.try_batch_graph_wgpu(&mut hid, &[0], 1, None),
20111            crate::gpu::BatchGraphOutcome::Declined
20112        );
20113        assert_eq!(hid, emb, "a declined batch graph leaves the rows untouched");
20114        assert!(p.try_multi_burst(3, 0, 4).is_none());
20115        assert!(
20116            p.graph_declines().is_empty(),
20117            "no attention decline logged: {:?}",
20118            p.graph_declines()
20119        );
20120        // (No assertion on graph_prefill_preferred: with no attention
20121        // decline it follows the device — a test process that brought a
20122        // wgpu adapter up routes this resident MoE through the graph.)
20123
20124        let plain = || create_test_pipeline(8, 16, 2, 1, 4, 2, 32);
20125        assert_eq!(plain().graph_attn_decline_reason(), None);
20126        assert_eq!(plain().wgpu_graph_attn_decline(), None);
20127        assert!(
20128            plain().graph_attn_geom(0).is_none(),
20129            "uniform models keep the historical arms"
20130        );
20131        let mut q = plain();
20132        q.set_layer_sinks(1, vec![0.25, -0.25]).unwrap();
20133        assert_eq!(
20134            q.graph_attn_decline_reason(),
20135            Some("learned attention sinks")
20136        );
20137        assert_eq!(
20138            q.graph_attn_geom(1).unwrap().sink,
20139            Some(&[0.25f32, -0.25][..])
20140        );
20141        let mut q = plain();
20142        q.set_attn_geometry(None, Some(2)).unwrap();
20143        assert_eq!(
20144            q.graph_attn_decline_reason(),
20145            Some("V heads narrower than Q/K heads")
20146        );
20147        assert_eq!(q.graph_attn_geom(0).unwrap().dv, 2);
20148        let mut q = plain();
20149        q.sliding_layers = Some(vec![true, false]);
20150        q.swa = Some((4, usize::MAX));
20151        assert_eq!(q.graph_attn_decline_reason(), Some("sliding-window layers"));
20152        assert_eq!(q.graph_attn_geom(0).unwrap().window, Some(4));
20153        assert_eq!(q.graph_attn_geom(1).unwrap().window, None);
20154
20155        // Outside the per-layer geometry: a named wgpu decline, logged
20156        // once per site.
20157        let mut q = mimo_test_pipeline();
20158        q.rope_scale = 2.0;
20159        assert_eq!(
20160            q.wgpu_graph_attn_decline(),
20161            Some("scaled RoPE positions with per-layer geometry")
20162        );
20163        let emb = q.embed_single(3);
20164        assert!(
20165            q.try_token_graph_wgpu_steps(&emb, 0, &mut lg, 1, None, None, 0, q.num_layers)
20166                .is_none()
20167        );
20168        let _ = q.try_token_graph_wgpu_steps(&emb, 1, &mut lg, 1, None, None, 0, q.num_layers);
20169        let lines = q.graph_declines();
20170        assert_eq!(
20171            lines
20172                .iter()
20173                .filter(|l| l.starts_with("wgpu token graph") && l.contains("scaled RoPE"))
20174                .count(),
20175            1,
20176            "{lines:?}"
20177        );
20178    }
20179
20180    #[test]
20181    fn mimo_verify_rewind_preserves_lagging_host_caches() {
20182        let mut p = mimo_test_pipeline();
20183        for (li, layer) in p.kv_cache.layers.iter_mut().enumerate() {
20184            let row = vec![0.0; layer.num_kv_heads * layer.head_dim];
20185            for _ in 0..if li == 0 { 2 } else { 12 } {
20186                layer.append(&row, &row, &[]);
20187            }
20188        }
20189        p.mimo_verify_rewind(9).unwrap();
20190        assert_eq!(p.kv_cache.layers[0].seq_len, 2);
20191        for layer in &p.kv_cache.layers[1..] {
20192            assert_eq!(layer.seq_len, 9);
20193        }
20194    }
20195
20196    /// CMF_LAYER_DUMP: the decode walk and the batched prefill both write
20197    /// every (position, layer) hidden, the two sets agree byte for byte,
20198    /// and the last layer's file is the stack output.
20199    #[test]
20200    fn layer_dump_covers_every_position_and_layer_on_both_walks() {
20201        let dir = std::env::temp_dir().join(format!("cmf-layer-dump-{}", std::process::id()));
20202        let _ = std::fs::remove_dir_all(&dir);
20203        let mut p = mimo_test_pipeline();
20204        let hs = p.hidden_size;
20205        let ids = [5u32, 9, 11, 2, 40];
20206        p.layer_dump = Some(dir.join("decode"));
20207        for (pos, &id) in ids.iter().enumerate() {
20208            let e = p.embed_single(id);
20209            let _ = p.forward_layers(&e, pos, None);
20210        }
20211        p.clear_sequence_state();
20212        p.layer_dump = Some(dir.join("prefill"));
20213        let hb = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
20214        for pos in 0..ids.len() {
20215            for li in 0..p.num_layers {
20216                let name = format!("p{pos:06}_l{li:02}.f32");
20217                let a = std::fs::read(dir.join("decode").join(&name)).unwrap();
20218                let b = std::fs::read(dir.join("prefill").join(&name)).unwrap();
20219                assert_eq!(a.len(), hs * 4, "{name}");
20220                assert_eq!(a, b, "{name}");
20221            }
20222        }
20223        let last = std::fs::read(dir.join("prefill").join("p000004_l03.f32")).unwrap();
20224        let vals: Vec<f32> = last
20225            .chunks(4)
20226            .map(|c| f32::from_le_bytes(c.try_into().unwrap()))
20227            .collect();
20228        assert_eq!(f32_bits(&vals), f32_bits(&hb[4 * hs..5 * hs]));
20229        let _ = std::fs::remove_dir_all(&dir);
20230    }
20231
20232    #[test]
20233    fn attn_geometry_and_sinks_are_validated() {
20234        let mut p = create_test_pipeline(8, 16, 4, 2, 4, 2, 32);
20235        assert!(
20236            p.set_attn_geometry(Some(vec![2]), None).is_err(),
20237            "one entry per layer"
20238        );
20239        assert!(
20240            p.set_attn_geometry(Some(vec![2, 3]), None).is_err(),
20241            "3 does not divide 4"
20242        );
20243        assert!(p.set_attn_geometry(Some(vec![2, 0]), None).is_err());
20244        assert!(p.set_attn_geometry(None, Some(0)).is_err());
20245        assert!(
20246            p.set_attn_geometry(None, Some(5)).is_err(),
20247            "V wider than the head"
20248        );
20249        p.set_attn_geometry(None, Some(4)).unwrap();
20250        assert_eq!(
20251            p.v_head_dim, None,
20252            "v_head_dim == head_dim is the uniform case"
20253        );
20254        p.set_layer_sinks(1, vec![0.1; 4]).unwrap();
20255        p.set_attn_geometry(Some(vec![1, 4]), None).unwrap();
20256        assert_eq!(p.kv_cache.layers[0].num_kv_heads, 1);
20257        assert_eq!(p.kv_cache.layers[1].num_kv_heads, 4);
20258        assert!(
20259            p.kv_cache.layers[1].sinks.is_some(),
20260            "a reshape keeps the layer's sinks"
20261        );
20262        assert_eq!(p.layer_geom(1).0, 4);
20263        assert!(
20264            p.set_layer_sinks(0, vec![0.0; 3]).is_err(),
20265            "one sink per Q head"
20266        );
20267        assert!(p.set_layer_sinks(7, vec![0.0; 4]).is_err());
20268        assert!(p.set_layer_sinks(0, vec![f32::NAN, 0.0, 0.0, 0.0]).is_err());
20269    }
20270
20271    /// The O(1) Nyström state replaces a plain full-context softmax; it
20272    /// must never be armed on a sliding, sink or narrow-V layer.
20273    #[test]
20274    fn o1_is_never_armed_on_sink_window_or_narrow_v_layers() {
20275        let cfg = || {
20276            Some(crate::nystrom::O1Cfg {
20277                layers: crate::nystrom::O1Layers::All,
20278                m: 4,
20279                w: 8,
20280                sink: 2,
20281                rect: crate::nystrom::O1Rect::Aggregate,
20282            })
20283        };
20284        let mut p = mimo_test_pipeline();
20285        p.set_o1(cfg());
20286        assert!(!p.o1_active(), "every MiMo-shaped layer is ineligible");
20287        let mut q = create_test_pipeline(8, 16, 2, 1, 4, 3, 64);
20288        q.set_layer_sinks(1, vec![0.0, 0.0]).unwrap();
20289        q.sliding_layers = Some(vec![false, false, true]);
20290        q.swa = Some((4, usize::MAX));
20291        q.set_o1(cfg());
20292        assert_eq!(q.o1_flags, vec![true, false, false]);
20293    }
20294
20295    #[test]
20296    fn o1_batch_transition_publishes_one_epoch_before_serial_handoff() {
20297        const B: usize = 19;
20298        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
20299        p.set_o1(Some(crate::nystrom::O1Cfg {
20300            layers: crate::nystrom::O1Layers::All,
20301            m: 4,
20302            w: 8,
20303            sink: 2,
20304            rect: crate::nystrom::O1Rect::Aggregate,
20305        }));
20306        p.o1_begin_with_prefix(Some(B));
20307        let ids: Vec<u32> = (0..B as u32).collect();
20308        let _ = p.prefill_batch_span(PrefillIn::Ids(&ids), 0, None, 0, p.num_layers);
20309
20310        assert_eq!(p.o1_epoch, 1, "all layers publish one completed transition");
20311        assert!(p.kv_cache.layers.iter().all(|l| l.o1_sealed()));
20312        let next = p.embed_single(B as u32);
20313        let _ = p.forward_layers(&next, B, None);
20314        assert_eq!(p.o1_epoch, 1, "sealed handoff must not republish the epoch");
20315    }
20316
20317    #[test]
20318    fn o1_pair_transition_commits_scratch_before_epoch_publication() {
20319        const B: usize = 19;
20320        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 2, 260);
20321        // Keep a real recurrent layer ahead of the Full O(1) layer so the
20322        // pair test observes the GDN lane-2 scratch swap at the same
20323        // boundary, rather than only exercising an artificial scratch vec.
20324        let gdn_cfg = crate::linear_core::GdnCfg {
20325            num_v_heads: 2,
20326            num_k_heads: 1,
20327            key_head_dim: 2,
20328            value_head_dim: 4,
20329            conv_kernel: 3,
20330            hidden_size: 8,
20331            rms_eps: 1e-6,
20332            output_gate_sigmoid: false,
20333        };
20334        let synth = |n: usize, salt: usize| -> Vec<f32> {
20335            (0..n)
20336                .map(|i| (((i * 13 + salt * 7) % 97) as f32 / 97.0 - 0.5) * 0.4)
20337                .collect()
20338        };
20339        let qt = |rows: usize, cols: usize, salt: usize| {
20340            crate::qtensor::QTensor::from_f32(synth(rows * cols, salt), rows, cols)
20341        };
20342        let c_dim = gdn_cfg.conv_dim();
20343        let vd = gdn_cfg.num_v_heads * gdn_cfg.value_head_dim;
20344        p.weights.layers[0].attn = AttnKind::LinearGdn(crate::linear_core::GdnWeights {
20345            in_proj_qkv: qt(c_dim, 8, 1),
20346            in_proj_z: qt(vd, 8, 2),
20347            in_proj_a: qt(gdn_cfg.num_v_heads, 8, 3),
20348            in_proj_b: qt(gdn_cfg.num_v_heads, 8, 4),
20349            conv1d: synth(c_dim * gdn_cfg.conv_kernel, 5),
20350            a_log: vec![0.2, 0.5],
20351            dt_bias: synth(gdn_cfg.num_v_heads, 6),
20352            norm: vec![1.0; gdn_cfg.value_head_dim],
20353            out_proj: qt(8, vd, 7),
20354        });
20355        p.gdn_cfg = Some(gdn_cfg);
20356        p.set_o1(Some(crate::nystrom::O1Cfg {
20357            layers: crate::nystrom::O1Layers::All,
20358            m: 4,
20359            w: 8,
20360            sink: 2,
20361            rect: crate::nystrom::O1Rect::Aggregate,
20362        }));
20363        p.o1_begin_with_prefix(Some(B));
20364        for pos in 0..B - 2 {
20365            let emb = p.embed_single(pos as u32);
20366            let _ = p.forward_layers(&emb, pos, None);
20367        }
20368        let lane1_state = p.kv_cache.layers[0].linear_state.clone();
20369
20370        let e1 = p.embed_single((B - 2) as u32);
20371        let e2 = p.embed_single((B - 1) as u32);
20372        let _ = p.forward_pair(&e1, &e2, B - 2);
20373
20374        assert_eq!(p.o1_epoch, 1, "pair crossing B publishes one epoch");
20375        assert!(
20376            p.kv_cache
20377                .layers
20378                .iter()
20379                .enumerate()
20380                .all(|(li, l)| !p.o1_flags[li] || l.o1_sealed())
20381        );
20382        assert!(!p.kv_cache.layers[0].linear_state.is_empty());
20383        assert_ne!(
20384            p.kv_cache.layers[0].linear_state, lane1_state,
20385            "real pair must commit GDN lane 2 before returning"
20386        );
20387        assert!(p.kv_cache.layers[0].linear_scratch.is_empty());
20388        let next = p.embed_single(B as u32);
20389        let _ = p.forward_layers(&next, B, None);
20390        assert_eq!(p.o1_epoch, 1, "serial continuation must reuse the epoch");
20391    }
20392
20393    #[test]
20394    fn o1_error_observation_stays_terminal_until_reset() {
20395        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20396        p.set_o1(Some(crate::nystrom::O1Cfg {
20397            layers: crate::nystrom::O1Layers::All,
20398            m: 4,
20399            w: 8,
20400            sink: 2,
20401            rect: crate::nystrom::O1Rect::Aggregate,
20402        }));
20403        p.o1_begin();
20404        p.kv_cache.layers[0].o1_abort("synthetic transition failure".into());
20405
20406        assert!(p.o1_seal_checked().is_err());
20407        assert!(
20408            p.o1_seal_checked().is_err(),
20409            "retry must see the sticky error"
20410        );
20411        let k = vec![0.2f32; 4];
20412        let v = vec![0.3f32; 4];
20413        p.kv_cache.layers[0].append(&k, &v, &[]);
20414        assert_eq!(p.kv_cache.layers[0].seq_len, 0);
20415
20416        p.reset_session();
20417        p.o1_begin();
20418        p.kv_cache.layers[0].append(&k, &v, &[]);
20419        assert_eq!(p.kv_cache.layers[0].seq_len, 1);
20420    }
20421
20422    #[test]
20423    fn nll_graph_failure_is_terminal_and_request_is_reusable() {
20424        let ids = vec![1u32, 2, 3, 4, 5, 6];
20425        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20426        p.graph_logits = Some(vec![123.0]);
20427        p.graph_want_logits = true;
20428        p.graph_failed
20429            .store(true, std::sync::atomic::Ordering::Relaxed);
20430        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
20431        let err = p.nll_ids_from(&ids, 0).expect_err("prior graph failure");
20432        assert!(err.contains("before NLL"));
20433        assert!(p.graph_logits.is_none());
20434        assert!(!p.graph_want_logits);
20435        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20436        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20437
20438        let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20439        let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
20440        let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
20441        assert_eq!(actual.1, expected.1);
20442        assert!((actual.0 - expected.0).abs() < 1e-9);
20443    }
20444
20445    #[test]
20446    fn nll_forward_failure_discards_partial_score_and_clears_sidechannels() {
20447        let ids = vec![1u32, 2, 3, 4, 5, 6];
20448        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20449        p.nll_test_fail_at = Some(1);
20450        let err = p
20451            .nll_ids_from(&ids, 0)
20452            .expect_err("one-shot forward failure");
20453        assert!(err.contains("forward") || err.contains("score row"));
20454        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20455        assert!(!p.graph_want_logits);
20456        assert!(p.graph_logits.is_none());
20457        assert!(p.kv_history.is_empty());
20458
20459        let mut fresh = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20460        let expected = fresh.nll_ids_from(&ids, 0).expect("fresh NLL");
20461        let actual = p.nll_ids_from(&ids, 0).expect("reused NLL");
20462        assert_eq!(actual.1, expected.1);
20463        assert!((actual.0 - expected.0).abs() < 1e-9);
20464    }
20465
20466    #[test]
20467    fn nll_serial_failure_before_first_row_is_reported() {
20468        let ids = vec![1u32, 2, 3, 4];
20469        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20470        p.nll_test_force_serial = true;
20471        p.nll_test_fail_at = Some(0);
20472        let err = p.nll_ids_from(&ids, 0).expect_err("serial forward failure");
20473        assert!(err.contains("serial forward"));
20474        assert!(p.kv_history.is_empty());
20475        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20476        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20477    }
20478
20479    #[test]
20480    fn ffn_probe_failure_discards_recorder_and_state() {
20481        let ids = vec![1u32, 2, 3, 4];
20482        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20483        p.nll_test_fail_at = Some(0);
20484        let err = p
20485            .probe_ffn_mass_batch(&ids)
20486            .expect_err("probe forward failure");
20487        assert!(err.contains("NLL"));
20488        assert!(FFN_PROBE.with(|probe| probe.borrow().is_none()));
20489        assert!(p.kv_history.is_empty());
20490        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20491    }
20492
20493    #[test]
20494    fn nll_test_controls_are_pipeline_scoped() {
20495        let ids = vec![1u32, 2, 3, 4];
20496        let mut failing = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20497        let mut unaffected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20498        failing.nll_test_force_serial = true;
20499        failing.nll_test_fail_at = Some(0);
20500
20501        assert!(!failing.can_prefill_batched());
20502        assert!(unaffected.can_prefill_batched());
20503        let expected = unaffected
20504            .nll_ids_from(&ids, 0)
20505            .expect("unaffected pipeline remains usable");
20506        let err = failing
20507            .nll_ids_from(&ids, 0)
20508            .expect_err("failure injection belongs to failing pipeline");
20509        assert!(err.contains("serial forward"));
20510        assert!(failing.nll_test_fail_at.is_none());
20511        assert!(unaffected.can_prefill_batched());
20512        let actual = unaffected
20513            .nll_ids_from(&ids, 0)
20514            .expect("unaffected pipeline remains reusable");
20515        assert_eq!(actual.1, expected.1);
20516        assert!((actual.0 - expected.0).abs() < 1e-9);
20517    }
20518
20519    #[test]
20520    fn forward_ids_failure_channel_is_terminal_and_reusable() {
20521        let ids = vec![1u32, 2, 3, 4, 5, 6];
20522        let mut p = create_test_pipeline(8, 16, 2, 1, 4, 1, 64);
20523        p.graph_logits = Some(vec![123.0]);
20524        p.graph_want_logits = true;
20525        p.graph_failed
20526            .store(true, std::sync::atomic::Ordering::Relaxed);
20527        p.cancel.store(true, std::sync::atomic::Ordering::Relaxed);
20528
20529        let err = p
20530            .forward_ids(&ids, None)
20531            .expect_err("a failed forward must not become a valid head result");
20532        assert!(err.contains("forward_ids setup"));
20533        assert!(p.graph_logits.is_none());
20534        assert!(!p.graph_want_logits);
20535        assert!(!p.graph_failed.load(std::sync::atomic::Ordering::Relaxed));
20536        assert!(!p.cancel.load(std::sync::atomic::Ordering::Relaxed));
20537        assert_eq!(p.kv_cache.seq_len(), 0);
20538
20539        let expected = create_test_pipeline(8, 16, 2, 1, 4, 1, 64)
20540            .forward_ids(&ids, None)
20541            .expect("fresh forward_ids");
20542        let actual = p
20543            .forward_ids(&ids, None)
20544            .expect("pipeline remains reusable after a failed forward");
20545        assert_eq!(actual.len(), expected.len());
20546        assert!(
20547            actual
20548                .iter()
20549                .zip(expected)
20550                .all(|(a, b)| (a - b).abs() < 1e-9)
20551        );
20552        assert_eq!(p.kv_cache.seq_len(), ids.len());
20553    }
20554
20555    #[test]
20556    fn sigmoid_router_floor_is_explicit_per_architecture() {
20557        // GLM-5's noaux_tc reference uses +1e-20 while the generic
20558        // LFM2-compatible path uses +1e-6.  At low (but representable)
20559        // sigmoid scores, silently sharing the latter changes expert weights
20560        // by orders of magnitude and can make a routed layer look coherent
20561        // while discarding its expert contribution.
20562        let zero = || DenseFfn {
20563            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20564            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20565            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20566            act: Act::Silu,
20567            down_t: None,
20568            segs: Vec::new(),
20569        };
20570        let m = MoeFfn {
20571            router: QTensor::from_f32(vec![0.0; 4], 2, 2),
20572            experts: vec![zero(), zero()],
20573            top_k: 1,
20574            norm_topk_prob: true,
20575            router_sigmoid: true,
20576            expert_bias: None,
20577            routed_scaling: 2.5,
20578            route_tau: None,
20579            shared: None,
20580            stats: std::cell::RefCell::new(Vec::new()),
20581            act_sq: std::cell::RefCell::new(Vec::new()),
20582            act_rows: std::cell::RefCell::new(Vec::new()),
20583            mask: None,
20584            per_expert_scale: None,
20585            router_input_norm: false,
20586            resonance: None,
20587            grown: Vec::new(),
20588        };
20589        let logits = [-20.0f32, -20.0];
20590        let (_, p, glm_wsum) = moe_route_with_eps(&logits, &m, None, 1e-20);
20591        let (_, _, generic_wsum) = moe_route(&logits, &m, None);
20592        let expected = (p[0] + 1e-20) / m.routed_scaling;
20593        assert!((glm_wsum - expected).abs() < 1e-15);
20594        assert!(generic_wsum > glm_wsum * 100.0);
20595    }
20596
20597    #[test]
20598    fn resonance_scores_match_formula_and_stable_tie() {
20599        let r = Resonance {
20600            // Three descriptors, hidden=2, one projection row each.
20601            mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0],
20602            u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0],
20603            k: 1,
20604            bias: vec![1.5, 0.5, 0.0],
20605            shell: Vec::new(),
20606        };
20607        let x = [1.0f32, 1.0];
20608        let mut got = vec![0.0; 3];
20609        r.scores(&x, &mut got);
20610        // Expert 0 and 1 are an exact score tie; the CPU top-1 contract uses
20611        // the lower index.  The values also check d² - (U·d)², not just tie
20612        // ordering.
20613        assert!((got[0] - 0.5).abs() < 1e-6);
20614        assert!((got[1] - 0.5).abs() < 1e-6);
20615        assert!(got[2].abs() < 1e-6);
20616        let best = got
20617            .iter()
20618            .enumerate()
20619            .max_by(|(ia, a), (ib, b)| a.partial_cmp(b).unwrap().then(ib.cmp(ia)))
20620            .map(|(i, _)| i);
20621        assert_eq!(best, Some(0));
20622        assert!(got.iter().all(|v| v.is_finite()));
20623    }
20624
20625    /// The growth shell (spec §2): a grown expert whose reconstruction
20626    /// error lies outside its shell scores −∞, one inside keeps the exact
20627    /// resonance score, trunk rows (`+inf` shell) are bit-identical to the
20628    /// shell-less computation; the process-wide switch disables it.
20629    #[test]
20630    fn resonance_shell_masks_outside_keeps_inside_and_trunk_bits() {
20631        // hidden = 2, rank 1. Experts 0/1 = trunk (shell +inf); 2 and 3 =
20632        // grown, the same descriptor (μ = (0, 1), u = (1, 1)) with shells
20633        // 6.0 and 0.25. At x' = (3, 0): d = (3, −1), d² = 10, proj =
20634        // (3 − 1)² = 4, err = 6 exactly — on the boundary of expert 2's
20635        // shell (kept: the rule is strict `>`), outside expert 3's.
20636        let plain = Resonance {
20637            mu: vec![0.0, 0.0, 1.0, 0.0, 0.0, 1.0, 0.0, 1.0],
20638            u: vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0, 1.0, 1.0],
20639            k: 1,
20640            bias: vec![1.5, 0.5, 0.0, 0.0],
20641            shell: Vec::new(),
20642        };
20643        let shelled = Resonance {
20644            mu: plain.mu.clone(),
20645            u: plain.u.clone(),
20646            k: 1,
20647            bias: plain.bias.clone(),
20648            shell: vec![f32::INFINITY, f32::INFINITY, 6.0, 0.25],
20649        };
20650        assert!(!plain.has_shell());
20651        assert!(shelled.has_shell());
20652        let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
20653        let (mut a, mut b) = (vec![0.0; 4], vec![0.0; 4]);
20654        set_growth_shell(Some(true));
20655        assert!(growth_shell_enabled());
20656        // x = (1, 1): the grown experts reconstruct it exactly (err 0):
20657        // inside both shells, every row the shell-less bits.
20658        let x = [1.0f32, 1.0];
20659        plain.scores(&x, &mut a);
20660        shelled.scores(&x, &mut b);
20661        assert_eq!(bits(&a), bits(&b), "inside every shell: unchanged");
20662        assert!(a[2] == 0.0 && a[3] == 0.0);
20663        // x' = (3, 0): expert 3 → −∞, expert 2 (err == shell) and the
20664        // trunk rows keep their exact bits.
20665        let xo = [3.0f32, 0.0];
20666        plain.scores(&xo, &mut a);
20667        shelled.scores(&xo, &mut b);
20668        assert_eq!(a[2], -6.0);
20669        assert_eq!(a[3], -6.0);
20670        assert_eq!(bits(&a[..3]), bits(&b[..3]), "trunk rows + the boundary row");
20671        assert_eq!(b[3], f32::NEG_INFINITY, "outside the shell: −∞");
20672        assert_eq!(shelled.effective_shell(4), shelled.shell);
20673        // The switch (`CMF_GROWTH_SHELL=off` / `growth-eval --shell off`):
20674        // all +inf, the shell-less bits everywhere.
20675        set_growth_shell(Some(false));
20676        assert!(!growth_shell_enabled());
20677        assert_eq!(shelled.effective_shell(4), vec![f32::INFINITY; 4]);
20678        shelled.scores(&xo, &mut b);
20679        assert_eq!(bits(&a), bits(&b));
20680        set_growth_shell(None);
20681        // A shell vector shorter than the expert count masks nothing
20682        // beyond it (a legacy layer whose tail has no shell).
20683        let short = Resonance {
20684            shell: vec![f32::INFINITY, f32::INFINITY],
20685            ..shelled
20686        };
20687        set_growth_shell(Some(true));
20688        short.scores(&xo, &mut b);
20689        assert_eq!(bits(&a), bits(&b));
20690        set_growth_shell(None);
20691    }
20692
20693    /// `moe_route` with −∞ logits (a grown expert outside its shell):
20694    /// top-1 is the best finite expert with weight exactly 1.0 on both
20695    /// the softmax and the sigmoid path; all −∞ degrades to uniform.
20696    #[test]
20697    fn moe_route_neg_inf_logits_pick_the_best_finite_expert_with_weight_one() {
20698        let zero = || DenseFfn {
20699            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20700            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20701            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20702            act: Act::Silu,
20703            down_t: None,
20704            segs: Vec::new(),
20705        };
20706        let moe = |sigmoid: bool| MoeFfn {
20707            router: QTensor::from_f32(vec![0.0; 8], 4, 2),
20708            experts: vec![zero(), zero(), zero(), zero()],
20709            top_k: 1,
20710            norm_topk_prob: true,
20711            router_sigmoid: sigmoid,
20712            expert_bias: None,
20713            routed_scaling: 1.0,
20714            route_tau: None,
20715            shared: None,
20716            stats: std::cell::RefCell::new(Vec::new()),
20717            act_sq: std::cell::RefCell::new(Vec::new()),
20718            act_rows: std::cell::RefCell::new(Vec::new()),
20719            mask: None,
20720            per_expert_scale: None,
20721            router_input_norm: false,
20722            resonance: None,
20723            grown: Vec::new(),
20724        };
20725        let logits = [-1.0f32, f32::NEG_INFINITY, -0.5, f32::NEG_INFINITY];
20726        for sigmoid in [false, true] {
20727            let m = moe(sigmoid);
20728            let (idx, p, wsum) = moe_route(&logits, &m, None);
20729            assert_eq!(idx, vec![2], "sigmoid {sigmoid}");
20730            assert_eq!(p[1], 0.0);
20731            assert_eq!(p[3], 0.0);
20732            assert!(p[2] > p[0] && p[0] > 0.0);
20733            assert!(p.iter().all(|v| v.is_finite()));
20734            let w = p[2] / wsum;
20735            if sigmoid {
20736                // The sigmoid renorm keeps its reference floor (+1e-6).
20737                assert!((w - 1.0).abs() < 1e-5, "sigmoid: weight {w}");
20738            } else {
20739                assert_eq!(w, 1.0, "softmax: the renormalized top-1 weight is exactly 1");
20740            }
20741            // Masked experts stay masked even when they are the only ones
20742            // "admitted" by an allow-list that covers everything.
20743            let (idx, _, _) = moe_route(&logits, &m, Some(&[true, true, true, true]));
20744            assert_eq!(idx, vec![2]);
20745        }
20746        // A finite expert always beats −∞ whatever the bias / order.
20747        let m = moe(false);
20748        let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, f32::NEG_INFINITY, f32::NEG_INFINITY, -9.0], &m, None);
20749        assert_eq!(idx, vec![3]);
20750        // Every expert at −∞ (cannot happen on a grown file — trunk rows
20751        // have no shell): uniform, finite, lowest index.
20752        let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 4], &m, None);
20753        assert_eq!(idx, vec![0]);
20754        assert!(p.iter().all(|&v| v == 0.25));
20755        assert!(wsum.is_finite() && wsum > 0.0);
20756    }
20757
20758    /// The resonance router (top-1) selects by the raw score as the
20759    /// trainer and the graph do — not by softmax probabilities, where two
20760    /// scores closer than 2^-25 collapse to the same `exp(l − max) = 1.0`
20761    /// and the LOWER index wins a token whose score is strictly smaller.
20762    #[test]
20763    fn resonance_route_picks_the_raw_argmax_not_the_softmax_tie() {
20764        let zero = || DenseFfn {
20765            gate_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20766            up_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20767            down_proj: QTensor::from_f32(vec![0.0; 4], 2, 2),
20768            act: Act::Silu,
20769            down_t: None,
20770            segs: Vec::new(),
20771        };
20772        let moe = |resonant: bool, norm_topk: bool| MoeFfn {
20773            router: QTensor::from_f32(vec![0.0; 6], 3, 2),
20774            experts: vec![zero(), zero(), zero()],
20775            top_k: 1,
20776            norm_topk_prob: norm_topk,
20777            router_sigmoid: false,
20778            expert_bias: None,
20779            routed_scaling: 1.0,
20780            route_tau: None,
20781            shared: None,
20782            stats: std::cell::RefCell::new(Vec::new()),
20783            act_sq: std::cell::RefCell::new(Vec::new()),
20784            act_rows: std::cell::RefCell::new(Vec::new()),
20785            mask: None,
20786            per_expert_scale: None,
20787            router_input_norm: false,
20788            resonance: resonant.then(|| Resonance {
20789                mu: vec![0.0; 6],
20790                u: Vec::new(),
20791                k: 0,
20792                bias: vec![0.0; 3],
20793                shell: Vec::new(),
20794            }),
20795            grown: Vec::new(),
20796        };
20797        // lo = −0.1, hi = the next f32 towards zero: hi − lo = 2^-27 <
20798        // 2^-25, so exp(lo − hi) rounds to exactly 1.0 — a softmax tie.
20799        let lo = -0.1f32;
20800        let hi = f32::from_bits(lo.to_bits() - 1);
20801        assert!(hi > lo && hi - lo < 2f32.powi(-25));
20802        assert_eq!((lo - hi).exp(), 1.0, "the tie this test is about");
20803        // The gated MoE (softmax) path: the tie hands the token to index 0.
20804        let (idx, _, _) = moe_route(&[lo, hi], &moe(false, true), None);
20805        assert_eq!(idx, vec![0], "softmax tie → lower index (the gated contract)");
20806        // The resonance path: the strictly larger raw score wins, weight
20807        // exactly 1.0 with and without norm_topk.
20808        for norm in [true, false] {
20809            let m = moe(true, norm);
20810            let (idx, p, wsum) = moe_route(&[lo, hi, f32::NEG_INFINITY], &m, None);
20811            assert_eq!(idx, vec![1], "norm_topk {norm}");
20812            assert_eq!(p, vec![0.0, 1.0, 0.0]);
20813            assert_eq!(p[1] / wsum, 1.0);
20814            // An exact tie: the first maximum (as `resonance_winner` and
20815            // `embryo_core_route_pick`).
20816            let (idx, _, _) = moe_route(&[hi, hi, lo], &m, None);
20817            assert_eq!(idx, vec![0]);
20818            // `−∞` never wins; the admitted set is honoured.
20819            let (idx, _, _) = moe_route(&[f32::NEG_INFINITY, lo, hi], &m, None);
20820            assert_eq!(idx, vec![2]);
20821            let (idx, p, wsum) = moe_route(&[lo, hi, hi], &m, Some(&[true, false, false]));
20822            assert_eq!(idx, vec![0]);
20823            assert_eq!(p[0] / wsum, 1.0);
20824            // Every admitted expert at −∞: the generic path's uniform
20825            // fallback (lowest index, finite weights).
20826            let (idx, p, wsum) = moe_route(&[f32::NEG_INFINITY; 3], &m, None);
20827            assert_eq!(idx, vec![0]);
20828            assert!(p.iter().all(|v| v.is_finite()) && wsum.is_finite() && wsum > 0.0);
20829        }
20830    }
20831}