Skip to main content

memra_reference/
lib.rs

1//! Portable, deliberately unfused executor for canonical `ModelPlan` operations.
2//!
3//! This crate is a correctness oracle, not a serving backend. It has no CUDA dependency and no
4//! external engine fallback. Unsupported canonical operations return a named error.
5
6pub mod hidden_trace;
7
8use memra_gguf::config::AttentionGateKind;
9use memra_gguf::model_plan::{
10    ActivationPlan, AttentionPlan, AttentionScale, GdnGateActivation, GemmaLayerScale, HcCollapse,
11    LogitsTransform, MicroBlockIndexPlan, MlpPlan, ModelPlan, PleEmbeddingPlan, ResidualTopology,
12    RopePlan, ValueNorm, ValueProjection,
13};
14use memra_gguf::tensor_contract::{DsparkTensor, LayerTensor, MtpTensor, TensorId, VisionTensor};
15use std::collections::BTreeMap;
16
17#[derive(Debug, Clone, PartialEq)]
18pub struct ReferenceTensor {
19    /// Logical row-major shape, outermost dimension first.
20    pub shape: Vec<usize>,
21    pub data: Vec<f32>,
22    /// Exact payload for checkpoint I64 index tensors (qwen4_exp n-gram multipliers /
23    /// vocab sizes / offsets). f32 cannot carry them (the vocab primes >= 2e7 exceed the
24    /// 24-bit mantissa and the hash multipliers fill i64), so integer tensors keep `data`
25    /// empty and executors read them through `tensor_i64` only.
26    pub ints: Option<Vec<i64>>,
27}
28
29impl ReferenceTensor {
30    pub fn new(shape: Vec<usize>, data: Vec<f32>) -> Result<Self, ReferenceError> {
31        let expected = shape.iter().product();
32        if data.len() != expected {
33            return Err(ReferenceError::TensorShape {
34                id: None,
35                expected: shape,
36                actual_elements: data.len(),
37            });
38        }
39        Ok(Self {
40            shape,
41            data,
42            ints: None,
43        })
44    }
45
46    pub fn new_i64(shape: Vec<usize>, ints: Vec<i64>) -> Result<Self, ReferenceError> {
47        let expected = shape.iter().product();
48        if ints.len() != expected {
49            return Err(ReferenceError::TensorShape {
50                id: None,
51                expected: shape,
52                actual_elements: ints.len(),
53            });
54        }
55        Ok(Self {
56            shape,
57            data: Vec::new(),
58            ints: Some(ints),
59        })
60    }
61}
62
63pub type ReferenceWeights = BTreeMap<TensorId, ReferenceTensor>;
64
65#[derive(Debug, Clone, PartialEq)]
66pub struct ReferenceFixture {
67    pub token_ids: Vec<u32>,
68    pub weights: ReferenceWeights,
69    pub vision: Option<ReferenceVisionInput>,
70    pub multimodal_token_ids: Option<Vec<u32>>,
71}
72
73#[derive(Debug, Clone, PartialEq)]
74pub struct ReferenceVisionInput {
75    /// Patch rows, row-major `[patches, patch_row_width]`. The pixel contract is per
76    /// tower program:
77    /// - `VisionPlan::Factored` (gemma-4): raw pixels in `[0, 1]`, width
78    ///   `3 * patch_size^2`; the executor applies the graph's `2x - 1`.
79    /// - `VisionPlan::Glm5Fused`: PREPROCESSED pixels (rescale 1/255 then CLIP mean/std
80    ///   normalize, done by the image processor), width
81    ///   `in_channels * temporal_patch_size * patch_size^2` in `(c, t, ph, pw)` flat
82    ///   order, token sequence in spatial-merge block-major order.
83    pub patches: ReferenceTensor,
84    /// Per-patch 2D position. Factored: `[x, y]`. Glm5Fused: `[h, w]` (the upstream
85    /// `get_vision_position_ids` column order), block-major over merge blocks.
86    pub positions: Vec<[u32; 2]>,
87    pub output_tokens: usize,
88}
89
90#[derive(Debug, Clone, PartialEq)]
91pub struct ReferenceVisionOutput {
92    pub encoder_hidden: Vec<f32>,
93    pub pooled_hidden: Vec<f32>,
94    pub projected_hidden: Vec<f32>,
95    pub patch_count: usize,
96    pub output_tokens: usize,
97    pub hidden_size: usize,
98    pub projection_size: usize,
99}
100
101#[derive(Debug, Clone, PartialEq)]
102pub struct ReferenceMultimodalOutput {
103    pub language: ReferenceOutput,
104    pub vision: ReferenceVisionOutput,
105}
106
107#[derive(Debug, Clone, PartialEq)]
108pub struct ReferenceState {
109    pub layers: Vec<ReferenceLayerState>,
110}
111
112#[derive(Debug, Clone, PartialEq)]
113pub enum ReferenceLayerState {
114    Kv {
115        key: Vec<f32>,
116        value: Vec<f32>,
117        tokens: usize,
118        kv_heads: usize,
119        key_head_dim: usize,
120        value_head_dim: usize,
121        window: Option<usize>,
122    },
123    Recurrent {
124        conv: Vec<f32>,
125        matrix: Vec<f32>,
126        value_heads: usize,
127        key_head_dim: usize,
128        value_head_dim: usize,
129        conv_width: usize,
130    },
131    LatentKv {
132        rows: Vec<f32>,
133        tokens: usize,
134        width: usize,
135    },
136    CompressedAttention {
137        rows: Vec<f32>,
138        tokens: usize,
139        width: usize,
140        window: usize,
141        compressed_tokens: usize,
142    },
143}
144
145#[derive(Debug, Clone, PartialEq)]
146pub struct ReferenceOutput {
147    /// `[tokens, vocab]`, row-major.
148    pub logits: Vec<f32>,
149    pub tokens: usize,
150    pub vocab: usize,
151    pub state: ReferenceState,
152    pub mtp: Vec<ReferenceMtpOutput>,
153    pub draft: Option<ReferenceDraftOutput>,
154    /// Post-layer residual state per TRUNK layer, `[tokens, hidden]`, or the
155    /// WIDE `[tokens, streams * hidden]` stream for gated-residual (qwen4_exp)
156    /// and HyperConnections trunks. Parity-gate localization surface: layer i
157    /// here compares directly against a forward hook on decoder layer i of the
158    /// upstream Python implementation.
159    pub layer_hidden: Vec<Vec<f32>>,
160}
161
162#[derive(Debug, Clone, PartialEq)]
163pub struct ReferenceMtpOutput {
164    pub depth: u32,
165    pub logits: Vec<f32>,
166    /// Post-block hidden state, `[tokens, hidden]` — except gated-residual (qwen4_exp)
167    /// drafts, where it is the WIDE stream `[tokens, streams * hidden]`: the post-layer
168    /// wide state is the multi-step K>1 carrier (SEMANTICS.md §MTP).
169    pub hidden: Vec<f32>,
170    pub state: ReferenceLayerState,
171}
172
173#[derive(Debug, Clone, PartialEq)]
174pub struct ReferenceDraftOutput {
175    pub input_token: u32,
176    pub output_ids: Vec<u32>,
177    pub confidence: Vec<f32>,
178    pub logits: Vec<f32>,
179    pub hidden: Vec<f32>,
180    pub block_size: usize,
181}
182
183#[derive(Debug, Clone, PartialEq)]
184pub enum ReferenceError {
185    EmptyInput,
186    TokenOutOfRange {
187        token: u32,
188        vocab: usize,
189    },
190    MissingTensor(TensorId),
191    /// An I64 checkpoint tensor arrived without its exact integer payload — an f32 copy
192    /// would silently corrupt the n-gram hash arithmetic, so this refuses instead.
193    IntegerTensorRequired(TensorId),
194    TensorShape {
195        id: Option<TensorId>,
196        expected: Vec<usize>,
197        actual_elements: usize,
198    },
199    UnsupportedOperation {
200        layer: Option<u32>,
201        operation: &'static str,
202    },
203    InvalidPlan {
204        layer: Option<u32>,
205        reason: &'static str,
206    },
207}
208
209impl std::fmt::Display for ReferenceError {
210    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
211        match self {
212            Self::EmptyInput => write!(f, "reference executor requires at least one token"),
213            Self::TokenOutOfRange { token, vocab } => {
214                write!(f, "token {token} is outside vocabulary size {vocab}")
215            }
216            Self::MissingTensor(id) => write!(f, "missing reference tensor {id:?}"),
217            Self::IntegerTensorRequired(id) => {
218                write!(f, "reference tensor {id:?} must carry an exact I64 payload")
219            }
220            Self::TensorShape {
221                id,
222                expected,
223                actual_elements,
224            } => write!(
225                f,
226                "reference tensor {id:?} expected shape {expected:?}, got {actual_elements} elements"
227            ),
228            Self::UnsupportedOperation { layer, operation } => {
229                write!(
230                    f,
231                    "unsupported reference operation {operation} at layer {layer:?}"
232                )
233            }
234            Self::InvalidPlan { layer, reason } => {
235                write!(f, "invalid model plan at layer {layer:?}: {reason}")
236            }
237        }
238    }
239}
240
241impl std::error::Error for ReferenceError {}
242
243pub fn deterministic_fixture(plan: &ModelPlan) -> Result<ReferenceFixture, ReferenceError> {
244    let hidden = plan.hidden_size as usize;
245    let vocab = plan.vocab_size as usize;
246    if hidden == 0 || vocab < 2 || hidden > 256 || vocab > 262_144 {
247        return Err(ReferenceError::InvalidPlan {
248            layer: None,
249            reason: "reference fixture requires hidden<=256 and 2<=vocab<=262144",
250        });
251    }
252    let mut executable_layers: Vec<_> = plan
253        .layers
254        .iter()
255        .chain(plan.mtp_blocks.iter().map(|block| &block.layer))
256        .collect();
257    if let Some(memra_gguf::model_plan::DrafterPlan::Dspark(dspark)) = plan.drafter.as_ref() {
258        executable_layers.extend(dspark.blocks.iter());
259    }
260    let mut weights = ReferenceWeights::new();
261    weights.insert(
262        TensorId::TokenEmbedding,
263        generated_tensor(&[vocab, hidden], 1, 0.2)?,
264    );
265    let vision = match plan.vision.as_ref() {
266        Some(memra_gguf::model_plan::VisionPlan::Factored(vision)) => Some(add_vision_fixture(
267            &mut weights,
268            vision,
269            plan.multimodal
270                .and_then(|injection| injection.tokens_per_image),
271        )?),
272        Some(memra_gguf::model_plan::VisionPlan::Glm5Fused(vision)) => {
273            Some(add_vision_fixture_glm5(&mut weights, vision)?)
274        }
275        None => None,
276    };
277    if let Some(mixer) = plan.exit_mixer {
278        // qwen4_exp exits through the hyper_connection_mixer read gate — the census has NO
279        // final norm module (SEMANTICS.md §Layer stack), so no OutputNorm row here either.
280        add_exit_mixer_fixture(&mut weights, LayerScope::Trunk, &mixer, hidden, 240)?;
281        if !plan.mtp_blocks.is_empty() {
282            add_exit_mixer_fixture(
283                &mut weights,
284                LayerScope::Mtp { depth: 0 },
285                &mixer,
286                hidden,
287                245,
288            )?;
289        }
290    } else {
291        weights.insert(
292            TensorId::OutputNorm,
293            ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
294        );
295    }
296    let checkpoint_factor_width = executable_layers
297        .iter()
298        .copied()
299        .filter_map(|layer| match &layer.attention {
300            AttentionPlan::Full(attention) | AttentionPlan::SlidingWindow { attention, .. } => {
301                matches!(
302                    attention.rope.factors,
303                    memra_gguf::model_plan::RopeFactors::Checkpoint
304                )
305                .then_some(attention.rope.dimensions as usize / 2)
306            }
307            _ => None,
308        })
309        .max();
310    if let Some(width) = checkpoint_factor_width {
311        weights.insert(
312            TensorId::RopeFactors,
313            ReferenceTensor::new(vec![width], vec![1.0; width])?,
314        );
315    }
316    if let Some((streams, epsilon, sinkhorn_iterations, collapse)) = hyper_topology(plan)? {
317        // The mean collapse has no learned head tensors (Glm5NextTextHyperHead).
318        if collapse == HcCollapse::GatedHead {
319            add_hyper_head_fixture(&mut weights, streams, hidden)?;
320        }
321        if epsilon <= 0.0 || sinkhorn_iterations == 0 {
322            return Err(ReferenceError::InvalidPlan {
323                layer: None,
324                reason: "HyperConnections require positive epsilon and Sinkhorn iterations",
325            });
326        }
327    }
328    for layer in executable_layers {
329        match layer.residual {
330            ResidualTopology::Serial => {}
331            ResidualTopology::Gemma { parallel_moe, .. } => {
332                for tensor in [LayerTensor::PostAttentionNorm, LayerTensor::PostMlpNorm] {
333                    weights.insert(
334                        layer_id(layer.index, tensor),
335                        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
336                    );
337                }
338                weights.insert(
339                    layer_id(layer.index, LayerTensor::LayerScale),
340                    ReferenceTensor::new(vec![1], vec![0.9])?,
341                );
342                if parallel_moe.is_some() {
343                    for tensor in [
344                        LayerTensor::PostSharedMlpNorm,
345                        LayerTensor::PreRoutedMlpNorm,
346                        LayerTensor::PostRoutedMlpNorm,
347                    ] {
348                        weights.insert(
349                            layer_id(layer.index, tensor),
350                            ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
351                        );
352                    }
353                }
354            }
355            ResidualTopology::HyperConnections { streams, .. } => {
356                add_hyper_fixture(&mut weights, layer.index, streams as usize, hidden)?;
357            }
358            ResidualTopology::GatedResidual {
359                streams,
360                bottleneck_rank,
361            } => {
362                add_gated_residual_fixture(
363                    &mut weights,
364                    layer_scope(plan, layer.index),
365                    layer.index,
366                    streams as usize,
367                    bottleneck_rank as usize,
368                    hidden,
369                )?;
370            }
371        }
372        // qwen4_exp has no input_layernorm/post_attention_layernorm modules — the read
373        // gate's grouped hc_norm IS the sublayer normalization (SEMANTICS.md §Layer stack).
374        if !matches!(layer.residual, ResidualTopology::GatedResidual { .. }) {
375            for tensor in [LayerTensor::PreAttentionNorm, LayerTensor::PreMlpNorm] {
376                weights.insert(
377                    layer_id(layer.index, tensor),
378                    ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
379                );
380            }
381        }
382        if let Some(overlay) = layer.sparse_overlay.as_ref() {
383            add_micro_block_index_fixture(
384                &mut weights,
385                layer_scope(plan, layer.index),
386                layer.index,
387                overlay,
388                hidden,
389            )?;
390        }
391        if let Some(ple) = layer.ple.as_ref() {
392            let ResidualTopology::GatedResidual { streams, .. } = layer.residual else {
393                return Err(ReferenceError::InvalidPlan {
394                    layer: Some(layer.index),
395                    reason: "PLE fixtures require the gated-residual wide stream",
396                });
397            };
398            add_ple_fixture(
399                &mut weights,
400                layer_scope(plan, layer.index),
401                layer.index,
402                ple,
403                streams as usize,
404                hidden,
405            )?;
406        }
407        match &layer.attention {
408            AttentionPlan::Full(attention) | AttentionPlan::SlidingWindow { attention, .. } => {
409                add_full_attention_fixture(&mut weights, layer.index, attention, hidden)?;
410            }
411            AttentionPlan::GatedDeltaNet(gdn) => {
412                add_gdn_fixture(&mut weights, layer.index, gdn, hidden)?;
413            }
414            AttentionPlan::Mla(mla) => {
415                add_mla_fixture(&mut weights, layer.index, mla, hidden)?;
416            }
417            AttentionPlan::KimiDeltaNet(kda) => {
418                add_kda_fixture(&mut weights, layer.index, kda, hidden)?;
419            }
420        }
421        match &layer.mlp {
422            MlpPlan::Dense(mlp) => {
423                add_dense_mlp_fixture(&mut weights, layer.index, mlp, hidden)?;
424            }
425            MlpPlan::Moe(moe) => {
426                add_moe_fixture(&mut weights, layer.index, moe, hidden, vocab)?;
427                if matches!(
428                    layer.residual,
429                    ResidualTopology::Gemma {
430                        parallel_moe: Some(_),
431                        ..
432                    }
433                ) {
434                    add_gemma_parallel_moe_fixture(&mut weights, layer.index, moe, hidden)?;
435                }
436            }
437        }
438    }
439    for block in &plan.mtp_blocks {
440        match block.input.fusion {
441            memra_gguf::model_plan::MtpFusionPlan::ConcatenateProjection => {
442                for tensor in [MtpTensor::EmbeddingNorm, MtpTensor::HiddenNorm] {
443                    weights.insert(
444                        TensorId::Mtp {
445                            depth: block.depth,
446                            tensor,
447                        },
448                        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
449                    );
450                }
451                weights.insert(
452                    TensorId::Mtp {
453                        depth: block.depth,
454                        tensor: MtpTensor::FusionProjection,
455                    },
456                    generated_tensor(
457                        &[hidden, 2 * hidden],
458                        100 + block.depth as u64,
459                        1.0 / ((2 * hidden) as f32).sqrt(),
460                    )?,
461                );
462            }
463            // qwen4_exp: two separate projections; the hidden-side norm covers the WIDE
464            // stream (census mtp.pre_fc_norm_hidden [10240] — SEMANTICS.md §MTP).
465            memra_gguf::model_plan::MtpFusionPlan::SeparateProjections => {
466                let ResidualTopology::GatedResidual { streams, .. } = block.layer.residual else {
467                    return Err(ReferenceError::InvalidPlan {
468                        layer: Some(block.layer.index),
469                        reason: "separate-projection MTP fusion requires a gated-residual block",
470                    });
471                };
472                weights.insert(
473                    TensorId::Mtp {
474                        depth: block.depth,
475                        tensor: MtpTensor::EmbeddingNorm,
476                    },
477                    ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
478                );
479                weights.insert(
480                    TensorId::Mtp {
481                        depth: block.depth,
482                        tensor: MtpTensor::HiddenNorm,
483                    },
484                    ReferenceTensor::new(
485                        vec![streams as usize * hidden],
486                        vec![1.0; streams as usize * hidden],
487                    )?,
488                );
489                for (tensor, salt) in [
490                    (MtpTensor::EmbeddingProjection, 230),
491                    (MtpTensor::HiddenProjection, 231),
492                ] {
493                    weights.insert(
494                        TensorId::Mtp {
495                            depth: block.depth,
496                            tensor,
497                        },
498                        generated_tensor(
499                            &[hidden, hidden],
500                            salt + block.depth as u64,
501                            1.0 / (hidden as f32).sqrt(),
502                        )?,
503                    );
504                }
505            }
506        }
507    }
508    if let Some(memra_gguf::model_plan::DrafterPlan::Dspark(dspark)) = plan.drafter.as_ref() {
509        add_dspark_fixture(&mut weights, dspark, hidden, vocab)?;
510    }
511    let token_ids = (1..=3.min(vocab - 1)).map(|token| token as u32).collect();
512    let multimodal_token_ids = plan.multimodal.map(|injection| {
513        // Grid-derived injection (glm5_next) takes the fixture image's own token count;
514        // fixed injection (gemma-4) takes the config-declared count.
515        let per_image = injection
516            .tokens_per_image
517            .map(|count| count as usize)
518            .or(vision.as_ref().map(|vision| vision.output_tokens))
519            .unwrap_or(1);
520        let mut tokens = Vec::with_capacity(per_image + 4);
521        tokens.push(1);
522        tokens.extend(injection.start_token_id);
523        tokens.extend(std::iter::repeat_n(
524            injection.placeholder_token_id,
525            per_image,
526        ));
527        tokens.extend(injection.end_token_id);
528        tokens.push(if injection.placeholder_token_id == 2 {
529            3
530        } else {
531            2
532        });
533        tokens
534    });
535    Ok(ReferenceFixture {
536        token_ids,
537        weights,
538        vision,
539        multimodal_token_ids,
540    })
541}
542
543fn add_dspark_fixture(
544    weights: &mut ReferenceWeights,
545    plan: &memra_gguf::model_plan::DsparkPlan,
546    hidden: usize,
547    vocab: usize,
548) -> Result<(), ReferenceError> {
549    if plan.blocks.is_empty()
550        || plan.block_size == 0
551        || plan.markov_rank == 0
552        || plan.target_layer_ids.is_empty()
553        || plan.noise_token_id as usize >= vocab
554    {
555        return Err(ReferenceError::InvalidPlan {
556            layer: None,
557            reason: "DSpark fixture requires blocks, targets, rank, block size, and valid noise token",
558        });
559    }
560    let streams = match plan.blocks[0].residual {
561        ResidualTopology::HyperConnections { streams, .. } if streams > 0 => streams as usize,
562        _ => {
563            return Err(ReferenceError::InvalidPlan {
564                layer: Some(plan.blocks[0].index),
565                reason: "DSpark blocks require HyperConnections",
566            });
567        }
568    };
569    weights.insert(
570        TensorId::Dspark(DsparkTensor::MainProjection),
571        generated_tensor(
572            &[hidden, plan.target_layer_ids.len() * hidden],
573            140,
574            1.0 / ((plan.target_layer_ids.len() * hidden) as f32).sqrt(),
575        )?,
576    );
577    weights.insert(
578        TensorId::Dspark(DsparkTensor::MainNorm),
579        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
580    );
581    weights.insert(
582        TensorId::Dspark(DsparkTensor::OutputNorm),
583        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
584    );
585    let rank = plan.markov_rank as usize;
586    weights.insert(
587        TensorId::Dspark(DsparkTensor::MarkovEmbedding),
588        generated_tensor(&[vocab, rank], 141, 0.1)?,
589    );
590    weights.insert(
591        TensorId::Dspark(DsparkTensor::MarkovOutput),
592        generated_tensor(&[vocab, rank], 142, 0.1)?,
593    );
594    weights.insert(
595        TensorId::Dspark(DsparkTensor::ConfidenceProjection),
596        generated_tensor(&[1, hidden + rank], 143, 0.1)?,
597    );
598    weights.insert(
599        TensorId::Dspark(DsparkTensor::HeadHyperFunction),
600        generated_tensor(&[streams, streams * hidden], 144, 0.1)?,
601    );
602    weights.insert(
603        TensorId::Dspark(DsparkTensor::HeadHyperBase),
604        generated_tensor(&[streams], 145, 0.05)?,
605    );
606    weights.insert(
607        TensorId::Dspark(DsparkTensor::HeadHyperScale),
608        ReferenceTensor::new(vec![1], vec![0.2])?,
609    );
610    Ok(())
611}
612
613fn add_vision_fixture(
614    weights: &mut ReferenceWeights,
615    plan: &memra_gguf::model_plan::VisionEncoderPlan,
616    output_tokens: Option<u32>,
617) -> Result<ReferenceVisionInput, ReferenceError> {
618    let hidden = plan.hidden_size as usize;
619    let patch_width =
620        (plan.patch.channels * plan.patch.patch_size * plan.patch.patch_size) as usize;
621    let axes = plan.patch.position_axes as usize;
622    let positions = plan.patch.position_embedding_size as usize;
623    weights.insert(
624        TensorId::Vision {
625            layer: None,
626            tensor: VisionTensor::PatchProjection,
627        },
628        generated_tensor(
629            &[hidden, patch_width],
630            150,
631            1.0 / (patch_width as f32).sqrt(),
632        )?,
633    );
634    weights.insert(
635        TensorId::Vision {
636            layer: None,
637            tensor: VisionTensor::PositionEmbedding,
638        },
639        generated_tensor(&[axes, positions, hidden], 151, 0.05)?,
640    );
641    if plan.standardize {
642        weights.insert(
643            TensorId::Vision {
644                layer: None,
645                tensor: VisionTensor::StandardizeBias,
646            },
647            generated_tensor(&[hidden], 152, 0.05)?,
648        );
649        weights.insert(
650            TensorId::Vision {
651                layer: None,
652                tensor: VisionTensor::StandardizeScale,
653            },
654            ReferenceTensor::new(vec![hidden], vec![0.5; hidden])?,
655        );
656    }
657    weights.insert(
658        TensorId::Vision {
659            layer: None,
660            tensor: VisionTensor::OutputProjection,
661        },
662        generated_tensor(
663            &[plan.projection_output_size as usize, hidden],
664            153,
665            1.0 / (hidden as f32).sqrt(),
666        )?,
667    );
668    for layer in &plan.layers {
669        let layer_id = Some(layer.index);
670        for tensor in [
671            VisionTensor::InputNorm,
672            VisionTensor::PostAttentionNorm,
673            VisionTensor::PreMlpNorm,
674            VisionTensor::PostMlpNorm,
675        ] {
676            weights.insert(
677                TensorId::Vision {
678                    layer: layer_id,
679                    tensor,
680                },
681                ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
682            );
683        }
684        let query_width = (layer.attention.query_heads * layer.attention.head_dim) as usize;
685        let kv_width = (layer.attention.kv_heads * layer.attention.head_dim) as usize;
686        for (tensor, shape, input, salt) in [
687            (VisionTensor::Query, vec![query_width, hidden], hidden, 160),
688            (VisionTensor::Key, vec![kv_width, hidden], hidden, 161),
689            (VisionTensor::Value, vec![kv_width, hidden], hidden, 162),
690            (
691                VisionTensor::AttentionOutput,
692                vec![hidden, query_width],
693                query_width,
694                163,
695            ),
696            (
697                VisionTensor::MlpGate,
698                vec![layer.mlp.intermediate_size as usize, hidden],
699                hidden,
700                164,
701            ),
702            (
703                VisionTensor::MlpUp,
704                vec![layer.mlp.intermediate_size as usize, hidden],
705                hidden,
706                165,
707            ),
708            (
709                VisionTensor::MlpDown,
710                vec![hidden, layer.mlp.intermediate_size as usize],
711                layer.mlp.intermediate_size as usize,
712                166,
713            ),
714        ] {
715            weights.insert(
716                TensorId::Vision {
717                    layer: layer_id,
718                    tensor,
719                },
720                generated_tensor(
721                    &shape,
722                    salt + layer.index as u64 * 17,
723                    1.0 / (input as f32).sqrt(),
724                )?,
725            );
726        }
727        for tensor in [VisionTensor::QueryNorm, VisionTensor::KeyNorm] {
728            weights.insert(
729                TensorId::Vision {
730                    layer: layer_id,
731                    tensor,
732                },
733                ReferenceTensor::new(
734                    vec![layer.attention.head_dim as usize],
735                    vec![1.0; layer.attention.head_dim as usize],
736                )?,
737            );
738        }
739    }
740    let side = plan.pooling_kernel_size.max(1) as usize;
741    let output_tokens = output_tokens.unwrap_or(1) as usize;
742    let patch_count = side * side * output_tokens;
743    let mut patches = generated_tensor(&[patch_count, patch_width], 170, 0.5)?;
744    for value in &mut patches.data {
745        *value += 0.5;
746    }
747    let mut patch_positions = Vec::with_capacity(patch_count);
748    for y in 0..side {
749        for x in 0..side * output_tokens {
750            patch_positions.push([x as u32, y as u32]);
751        }
752    }
753    Ok(ReferenceVisionInput {
754        patches,
755        positions: patch_positions,
756        output_tokens,
757    })
758}
759
760/// Deterministic tiny fixture for the glm5_next tower program. A 2-wide x 1-tall grid of
761/// spatial-merge blocks (`n = 2 * merge^2` patches, 2 output tokens) exercises both rope
762/// axes, the block-major position order, the downsample block gather and the merger.
763fn add_vision_fixture_glm5(
764    weights: &mut ReferenceWeights,
765    plan: &memra_gguf::model_plan::Glm5VisionPlan,
766) -> Result<ReferenceVisionInput, ReferenceError> {
767    let hidden = plan.hidden_size as usize;
768    let head_dim = plan.head_dim as usize;
769    let ff = plan.intermediate_size as usize;
770    let out = plan.out_hidden_size as usize;
771    let proj_inter = plan.projection_intermediate_size as usize;
772    let merge = plan.spatial_merge_size as usize;
773    let patch_width = plan.patch_input_width as usize;
774    let id = |layer: Option<u32>, tensor| TensorId::Vision { layer, tensor };
775    weights.insert(id(None, VisionTensor::PatchProjection), {
776        let mut tensor = generated_tensor(
777            &[hidden, patch_width],
778            150,
779            1.0 / (patch_width as f32).sqrt(),
780        )?;
781        // Census truth is the 5-d conv shape; row-major layout is identical.
782        tensor.shape = vec![
783            hidden,
784            plan.in_channels as usize,
785            plan.temporal_patch_size as usize,
786            plan.patch_size as usize,
787            plan.patch_size as usize,
788        ];
789        tensor
790    });
791    weights.insert(
792        id(None, VisionTensor::PatchProjectionBias),
793        generated_tensor(&[hidden], 151, 0.05)?,
794    );
795    for layer in 0..plan.depth {
796        let l = Some(layer);
797        let salt = layer as u64 * 23;
798        for (tensor, shape, input, seed) in [
799            (
800                VisionTensor::FusedQkv,
801                vec![3 * hidden, hidden],
802                hidden,
803                250,
804            ),
805            (
806                VisionTensor::AttentionOutput,
807                vec![hidden, hidden],
808                hidden,
809                251,
810            ),
811            (VisionTensor::MlpGate, vec![ff, hidden], hidden, 252),
812            (VisionTensor::MlpUp, vec![ff, hidden], hidden, 253),
813            (VisionTensor::MlpDown, vec![hidden, ff], ff, 254),
814        ] {
815            weights.insert(
816                id(l, tensor),
817                generated_tensor(&shape, seed + salt, 1.0 / (input as f32).sqrt())?,
818            );
819        }
820        for (tensor, width, seed) in [
821            (VisionTensor::FusedQkvBias, 3 * hidden, 255),
822            (VisionTensor::AttentionOutputBias, hidden, 256),
823            (VisionTensor::MlpGateBias, ff, 257),
824            (VisionTensor::MlpUpBias, ff, 258),
825            (VisionTensor::MlpDownBias, hidden, 259),
826        ] {
827            weights.insert(
828                id(l, tensor),
829                generated_tensor(&[width], seed + salt, 0.02)?,
830            );
831        }
832        for (tensor, width) in [
833            (VisionTensor::InputNorm, hidden),
834            (VisionTensor::PreMlpNorm, hidden),
835            (VisionTensor::QueryNorm, head_dim),
836            (VisionTensor::KeyNorm, head_dim),
837        ] {
838            weights.insert(
839                id(l, tensor),
840                ReferenceTensor::new(vec![width], vec![1.0; width])?,
841            );
842        }
843    }
844    weights.insert(
845        id(None, VisionTensor::PostEncoderNorm),
846        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
847    );
848    weights.insert(id(None, VisionTensor::Downsample), {
849        let mut tensor = generated_tensor(
850            &[out, hidden * merge * merge],
851            260,
852            1.0 / ((hidden * merge * merge) as f32).sqrt(),
853        )?;
854        tensor.shape = vec![out, hidden, merge, merge];
855        tensor
856    });
857    weights.insert(
858        id(None, VisionTensor::DownsampleBias),
859        generated_tensor(&[out], 261, 0.02)?,
860    );
861    weights.insert(
862        id(None, VisionTensor::MergerProjection),
863        generated_tensor(&[out, out], 262, 1.0 / (out as f32).sqrt())?,
864    );
865    weights.insert(
866        id(None, VisionTensor::MergerPostProjectionNorm),
867        ReferenceTensor::new(vec![out], vec![1.0; out])?,
868    );
869    weights.insert(
870        id(None, VisionTensor::MergerPostProjectionNormBias),
871        generated_tensor(&[out], 263, 0.02)?,
872    );
873    weights.insert(
874        id(None, VisionTensor::MergerGate),
875        generated_tensor(&[proj_inter, out], 264, 1.0 / (out as f32).sqrt())?,
876    );
877    weights.insert(
878        id(None, VisionTensor::MergerUp),
879        generated_tensor(&[proj_inter, out], 265, 1.0 / (out as f32).sqrt())?,
880    );
881    weights.insert(
882        id(None, VisionTensor::MergerDown),
883        generated_tensor(&[out, proj_inter], 266, 1.0 / (proj_inter as f32).sqrt())?,
884    );
885    // 1 x 2 blocks of merge x merge patches, block-major (block_row, block_col, in_row,
886    // in_col) — the upstream patchify/pos-id order.
887    let output_tokens = 2usize;
888    let patch_count = output_tokens * merge * merge;
889    let mut patches = generated_tensor(&[patch_count, patch_width], 270, 0.5)?;
890    for value in &mut patches.data {
891        *value += 0.5;
892    }
893    let mut positions = Vec::with_capacity(patch_count);
894    for block_col in 0..output_tokens {
895        for in_row in 0..merge {
896            for in_col in 0..merge {
897                positions.push([in_row as u32, (block_col * merge + in_col) as u32]);
898            }
899        }
900    }
901    Ok(ReferenceVisionInput {
902        patches,
903        positions,
904        output_tokens,
905    })
906}
907
908fn add_hyper_head_fixture(
909    weights: &mut ReferenceWeights,
910    streams: usize,
911    hidden: usize,
912) -> Result<(), ReferenceError> {
913    if streams == 0 {
914        return Err(ReferenceError::InvalidPlan {
915            layer: None,
916            reason: "HyperConnections require at least one stream",
917        });
918    }
919    weights.insert(
920        TensorId::HyperHeadFunction,
921        generated_tensor(&[streams, streams * hidden], 90, 0.1)?,
922    );
923    weights.insert(
924        TensorId::HyperHeadBase,
925        generated_tensor(&[streams], 91, 0.05)?,
926    );
927    weights.insert(
928        TensorId::HyperHeadScale,
929        ReferenceTensor::new(vec![1], vec![0.2])?,
930    );
931    Ok(())
932}
933
934fn add_hyper_fixture(
935    weights: &mut ReferenceWeights,
936    layer: u32,
937    streams: usize,
938    hidden: usize,
939) -> Result<(), ReferenceError> {
940    if streams == 0 {
941        return Err(ReferenceError::InvalidPlan {
942            layer: Some(layer),
943            reason: "HyperConnections require at least one stream",
944        });
945    }
946    let rows = (2 + streams) * streams;
947    for (function, base, scale, salt) in [
948        (
949            LayerTensor::HyperAttentionFunction,
950            LayerTensor::HyperAttentionBase,
951            LayerTensor::HyperAttentionScale,
952            92,
953        ),
954        (
955            LayerTensor::HyperMlpFunction,
956            LayerTensor::HyperMlpBase,
957            LayerTensor::HyperMlpScale,
958            95,
959        ),
960    ] {
961        weights.insert(
962            layer_id(layer, function),
963            generated_tensor(&[rows, streams * hidden], salt + layer as u64 * 101, 0.1)?,
964        );
965        weights.insert(
966            layer_id(layer, base),
967            generated_tensor(&[rows], salt + 1 + layer as u64 * 101, 0.05)?,
968        );
969        weights.insert(
970            layer_id(layer, scale),
971            ReferenceTensor::new(vec![3], vec![0.2, 0.2, 0.2])?,
972        );
973    }
974    Ok(())
975}
976
977/// Which checkpoint namespace a qwen4_exp layer's family-bound tensors live in. The pack
978/// binds gated-residual / indexer / PLE / mixer tensors as `TensorId::Family` keyed by
979/// `semantic_key` (crates/memra-gguf/src/model_packs/qwen4_exp: trunk rows strip the
980/// `model.language_model.` wrapper to `trunk.*`, MTP rows keep their `mtp.*` names). The
981/// key formats below MUST mirror that mapping — drift fails loudly as MissingTensor.
982#[derive(Debug, Clone, Copy, PartialEq, Eq)]
983enum LayerScope {
984    Trunk,
985    Mtp { depth: u32 },
986}
987
988impl LayerScope {
989    fn layer_prefix(self, index: u32) -> String {
990        match self {
991            Self::Trunk => format!("trunk.layers.{index}."),
992            Self::Mtp { depth } => format!("mtp.layers.{depth}."),
993        }
994    }
995
996    fn mixer_prefix(self) -> &'static str {
997        match self {
998            Self::Trunk => "trunk.hyper_connection_mixer.",
999            Self::Mtp { .. } => "mtp.hyper_connection_mixer.",
1000        }
1001    }
1002}
1003
1004/// Scope for a layer taken from `deterministic_fixture`'s combined trunk+MTP walk: the
1005/// plan compiles MTP block `depth` at global index `trunk_len + depth`.
1006fn layer_scope(plan: &ModelPlan, index: u32) -> LayerScope {
1007    let trunk = plan.layers.len() as u32;
1008    if index < trunk {
1009        LayerScope::Trunk
1010    } else {
1011        LayerScope::Mtp {
1012            depth: index - trunk,
1013        }
1014    }
1015}
1016
1017fn qwen4exp_family_id(key: String) -> TensorId {
1018    TensorId::Family {
1019        family: "qwen4_exp",
1020        key,
1021    }
1022}
1023
1024fn add_gated_residual_fixture(
1025    weights: &mut ReferenceWeights,
1026    scope: LayerScope,
1027    layer: u32,
1028    streams: usize,
1029    rank: usize,
1030    hidden: usize,
1031) -> Result<(), ReferenceError> {
1032    if streams == 0 || rank == 0 {
1033        return Err(ReferenceError::InvalidPlan {
1034            layer: Some(layer),
1035            reason: "gated residual requires streams and a bottleneck rank",
1036        });
1037    }
1038    let wide = streams * hidden;
1039    let prefix = scope.layer_prefix(layer);
1040    for (sublayer, salt) in [
1041        ("attn_hyper_connection.", 200u64),
1042        ("mlp_hyper_connection.", 204),
1043    ] {
1044        weights.insert(
1045            qwen4exp_family_id(format!("{prefix}{sublayer}hc_norm.weight")),
1046            ReferenceTensor::new(vec![wide], vec![1.0; wide])?,
1047        );
1048        weights.insert(
1049            qwen4exp_family_id(format!("{prefix}{sublayer}input_mix_weight_down.weight")),
1050            generated_tensor(
1051                &[rank, wide],
1052                salt + 1 + layer as u64 * 211,
1053                1.0 / (wide as f32).sqrt(),
1054            )?,
1055        );
1056        weights.insert(
1057            qwen4exp_family_id(format!("{prefix}{sublayer}input_mix_weight_up.weight")),
1058            generated_tensor(
1059                &[wide, rank],
1060                salt + 2 + layer as u64 * 211,
1061                1.0 / (rank as f32).sqrt(),
1062            )?,
1063        );
1064        weights.insert(
1065            qwen4exp_family_id(format!("{prefix}{sublayer}block_inject_weight.weight")),
1066            generated_tensor(
1067                &[streams, wide],
1068                salt + 3 + layer as u64 * 211,
1069                1.0 / (wide as f32).sqrt(),
1070            )?,
1071        );
1072    }
1073    Ok(())
1074}
1075
1076fn add_exit_mixer_fixture(
1077    weights: &mut ReferenceWeights,
1078    scope: LayerScope,
1079    mixer: &memra_gguf::model_plan::GatedResidualMixerPlan,
1080    hidden: usize,
1081    salt: u64,
1082) -> Result<(), ReferenceError> {
1083    let streams = mixer.streams as usize;
1084    let rank = mixer.bottleneck_rank as usize;
1085    if streams == 0 || rank == 0 {
1086        return Err(ReferenceError::InvalidPlan {
1087            layer: None,
1088            reason: "exit mixer requires streams and a bottleneck rank",
1089        });
1090    }
1091    let wide = streams * hidden;
1092    let prefix = scope.mixer_prefix();
1093    // Read half only: the census mixer carries NO block_inject row (use_combine=false).
1094    weights.insert(
1095        qwen4exp_family_id(format!("{prefix}hc_norm.weight")),
1096        ReferenceTensor::new(vec![wide], vec![1.0; wide])?,
1097    );
1098    weights.insert(
1099        qwen4exp_family_id(format!("{prefix}input_mix_weight_down.weight")),
1100        generated_tensor(&[rank, wide], salt + 1, 1.0 / (wide as f32).sqrt())?,
1101    );
1102    weights.insert(
1103        qwen4exp_family_id(format!("{prefix}input_mix_weight_up.weight")),
1104        generated_tensor(&[wide, rank], salt + 2, 1.0 / (rank as f32).sqrt())?,
1105    );
1106    Ok(())
1107}
1108
1109fn add_micro_block_index_fixture(
1110    weights: &mut ReferenceWeights,
1111    scope: LayerScope,
1112    layer: u32,
1113    overlay: &MicroBlockIndexPlan,
1114    hidden: usize,
1115) -> Result<(), ReferenceError> {
1116    let heads = overlay.query_heads as usize;
1117    let kv_heads = overlay.kv_heads as usize;
1118    let head_dim = overlay.head_dim as usize;
1119    if heads == 0 || kv_heads == 0 || head_dim == 0 || overlay.block_size == 0 {
1120        return Err(ReferenceError::InvalidPlan {
1121            layer: Some(layer),
1122            reason: "micro-block indexer requires heads, head_dim, and a block size",
1123        });
1124    }
1125    let prefix = scope.layer_prefix(layer);
1126    weights.insert(
1127        qwen4exp_family_id(format!("{prefix}self_attn.indexer.index_qk_proj.weight")),
1128        generated_tensor(
1129            &[(heads + kv_heads) * head_dim, hidden],
1130            210 + layer as u64 * 211,
1131            1.0 / (hidden as f32).sqrt(),
1132        )?,
1133    );
1134    for norm in ["q_layernorm", "k_layernorm"] {
1135        weights.insert(
1136            qwen4exp_family_id(format!("{prefix}self_attn.indexer.{norm}.weight")),
1137            ReferenceTensor::new(vec![head_dim], vec![1.0; head_dim])?,
1138        );
1139    }
1140    Ok(())
1141}
1142
1143fn add_ple_fixture(
1144    weights: &mut ReferenceWeights,
1145    scope: LayerScope,
1146    layer: u32,
1147    ple: &PleEmbeddingPlan,
1148    streams: usize,
1149    hidden: usize,
1150) -> Result<(), ReferenceError> {
1151    let heads = ple.ngram_heads as usize;
1152    let head_dim = ple.head_embed_dim as usize;
1153    let embed_dim = ple.embed_dim as usize;
1154    let kernel = ple.conv_kernel as usize;
1155    let max_ngram = ple.max_ngram as usize;
1156    if heads == 0
1157        || head_dim == 0
1158        || kernel == 0
1159        || max_ngram < 2
1160        || embed_dim != heads * head_dim
1161        || !heads.is_multiple_of(max_ngram - 1)
1162    {
1163        return Err(ReferenceError::InvalidPlan {
1164            layer: Some(layer),
1165            reason: "PLE fixture requires consistent n-gram head geometry",
1166        });
1167    }
1168    let wide = streams * hidden;
1169    let prefix = scope.layer_prefix(layer);
1170    weights.insert(
1171        qwen4exp_family_id(format!("{prefix}ple.key_proj.weight")),
1172        generated_tensor(
1173            &[wide, embed_dim],
1174            215 + layer as u64 * 211,
1175            1.0 / (embed_dim as f32).sqrt(),
1176        )?,
1177    );
1178    weights.insert(
1179        qwen4exp_family_id(format!("{prefix}ple.value_proj.weight")),
1180        generated_tensor(
1181            &[hidden, embed_dim],
1182            216 + layer as u64 * 211,
1183            1.0 / (embed_dim as f32).sqrt(),
1184        )?,
1185    );
1186    for norm in ["norm_key", "norm_query", "norm_conv"] {
1187        weights.insert(
1188            qwen4exp_family_id(format!("{prefix}ple.{norm}.weight")),
1189            ReferenceTensor::new(vec![wide], vec![1.0; wide])?,
1190        );
1191    }
1192    // Checkpoint ships [wide, 1, kernel]; the reference consumes the squeezed depthwise
1193    // form like the GDN conv row.
1194    weights.insert(
1195        qwen4exp_family_id(format!("{prefix}ple.conv1d.weight")),
1196        generated_tensor(
1197            &[wide, kernel],
1198            217 + layer as u64 * 211,
1199            1.0 / (kernel as f32).sqrt(),
1200        )?,
1201    );
1202    // Synthetic I64 index buffers. Real checkpoints LOAD these (SEMANTICS.md §PLE — never
1203    // re-derived); the fixture only needs deterministic, structurally valid values: odd
1204    // multipliers, distinct per-head vocab sizes with prefix-sum offsets, and a table with
1205    // a few pad rows past the addressable range.
1206    let multipliers: Vec<i64> = (0..max_ngram)
1207        .map(|index| 1_000_003 + 2 * (layer as i64 * 97 + index as i64 * 31))
1208        .collect();
1209    let sizes: Vec<i64> = (0..heads).map(|head| 17 + 2 * head as i64).collect();
1210    let mut offsets = Vec::with_capacity(heads);
1211    let mut total = 0i64;
1212    for &size in &sizes {
1213        offsets.push(total);
1214        total += size;
1215    }
1216    weights.insert(
1217        qwen4exp_family_id(format!("{prefix}ple.ple_embedding.layer_multipliers")),
1218        ReferenceTensor::new_i64(vec![max_ngram], multipliers)?,
1219    );
1220    weights.insert(
1221        qwen4exp_family_id(format!("{prefix}ple.ple_embedding.ngram_heads_vocab_sizes")),
1222        ReferenceTensor::new_i64(vec![heads], sizes)?,
1223    );
1224    weights.insert(
1225        qwen4exp_family_id(format!("{prefix}ple.ple_embedding.ngram_heads_offsets")),
1226        ReferenceTensor::new_i64(vec![heads], offsets)?,
1227    );
1228    weights.insert(
1229        qwen4exp_family_id(format!("{prefix}ple.ple_embedding.ngram_embedding")),
1230        generated_tensor(
1231            &[total as usize + 3, head_dim],
1232            218 + layer as u64 * 211,
1233            0.2,
1234        )?,
1235    );
1236    Ok(())
1237}
1238
1239#[allow(clippy::unusual_byte_groupings)] // allow: mnemonic grouping of a pinned seed/magic constant
1240fn generated_tensor(
1241    shape: &[usize],
1242    salt: u64,
1243    scale: f32,
1244) -> Result<ReferenceTensor, ReferenceError> {
1245    let elements = shape.iter().product();
1246    let data = (0..elements)
1247        .map(|index| {
1248            let mut value = index as u64 ^ salt.wrapping_mul(0x9e37_79b9);
1249            value ^= value >> 16;
1250            value = value.wrapping_mul(0x45d9_f3b);
1251            value ^= value >> 16;
1252            let unit = (value as u32) as f32 / u32::MAX as f32;
1253            (2.0 * unit - 1.0) * scale
1254        })
1255        .collect();
1256    ReferenceTensor::new(shape.to_vec(), data)
1257}
1258
1259fn add_full_attention_fixture(
1260    weights: &mut ReferenceWeights,
1261    layer: u32,
1262    attention: &memra_gguf::model_plan::FullAttentionPlan,
1263    hidden: usize,
1264) -> Result<(), ReferenceError> {
1265    let query_heads = attention.query_heads as usize;
1266    let kv_heads = attention.kv_heads as usize;
1267    let key_dim = attention.key_head_dim as usize;
1268    let value_dim = attention.value_head_dim as usize;
1269    let q_width = query_heads
1270        * key_dim
1271        * if attention.output_gate == AttentionGateKind::FusedQ {
1272            2
1273        } else {
1274            1
1275        };
1276    for (tensor, output, input, salt) in [
1277        (LayerTensor::Query, q_width, hidden, 10),
1278        (LayerTensor::Key, kv_heads * key_dim, hidden, 11),
1279        (
1280            LayerTensor::AttentionOutput,
1281            hidden,
1282            query_heads * value_dim,
1283            13,
1284        ),
1285    ] {
1286        weights.insert(
1287            layer_id(layer, tensor),
1288            generated_tensor(
1289                &[output, input],
1290                salt + layer as u64 * 31,
1291                1.0 / (input as f32).sqrt(),
1292            )?,
1293        );
1294    }
1295    if attention.value_projection == ValueProjection::Separate {
1296        weights.insert(
1297            layer_id(layer, LayerTensor::Value),
1298            generated_tensor(
1299                &[kv_heads * value_dim, hidden],
1300                12 + layer as u64 * 31,
1301                1.0 / (hidden as f32).sqrt(),
1302            )?,
1303        );
1304    }
1305    if attention.qk_norm != memra_gguf::model_plan::TensorPresence::Absent {
1306        for tensor in [LayerTensor::QueryNorm, LayerTensor::KeyNorm] {
1307            weights.insert(
1308                layer_id(layer, tensor),
1309                ReferenceTensor::new(vec![key_dim], vec![1.0; key_dim])?,
1310            );
1311        }
1312    }
1313    if attention.output_gate == AttentionGateKind::SeparateHead {
1314        weights.insert(
1315            layer_id(layer, LayerTensor::AttentionGate),
1316            generated_tensor(
1317                &[query_heads, hidden],
1318                14 + layer as u64 * 31,
1319                1.0 / (hidden as f32).sqrt(),
1320            )?,
1321        );
1322    }
1323    Ok(())
1324}
1325
1326fn add_gdn_fixture(
1327    weights: &mut ReferenceWeights,
1328    layer: u32,
1329    gdn: &memra_gguf::model_plan::GatedDeltaNetPlan,
1330    hidden: usize,
1331) -> Result<(), ReferenceError> {
1332    let key_heads = gdn.key_heads as usize;
1333    let value_heads = gdn.value_heads as usize;
1334    let key_dim = gdn.key_head_dim as usize;
1335    let value_dim = gdn.value_head_dim as usize;
1336    let conv_width = 2 * key_heads * key_dim + value_heads * value_dim;
1337    for (tensor, output, input, salt) in [
1338        (LayerTensor::GdnQkv, conv_width, hidden, 40),
1339        (LayerTensor::GdnGate, value_heads * value_dim, hidden, 41),
1340        (LayerTensor::GdnBeta, value_heads, hidden, 42),
1341        (LayerTensor::GdnAlpha, value_heads, hidden, 43),
1342        (LayerTensor::GdnOutput, hidden, value_heads * value_dim, 44),
1343    ] {
1344        weights.insert(
1345            layer_id(layer, tensor),
1346            generated_tensor(
1347                &[output, input],
1348                salt + layer as u64 * 47,
1349                1.0 / (input as f32).sqrt(),
1350            )?,
1351        );
1352    }
1353    weights.insert(
1354        layer_id(layer, LayerTensor::GdnA),
1355        ReferenceTensor::new(vec![value_heads], vec![-0.5; value_heads])?,
1356    );
1357    weights.insert(
1358        layer_id(layer, LayerTensor::GdnDtBias),
1359        ReferenceTensor::new(vec![value_heads], vec![0.0; value_heads])?,
1360    );
1361    weights.insert(
1362        layer_id(layer, LayerTensor::GdnNorm),
1363        ReferenceTensor::new(vec![value_dim], vec![1.0; value_dim])?,
1364    );
1365    weights.insert(
1366        layer_id(layer, LayerTensor::GdnConv1d),
1367        generated_tensor(
1368            &[conv_width, gdn.conv_kernel as usize],
1369            45 + layer as u64 * 47,
1370            1.0 / (gdn.conv_kernel as f32).sqrt(),
1371        )?,
1372    );
1373    Ok(())
1374}
1375
1376fn add_kda_fixture(
1377    weights: &mut ReferenceWeights,
1378    layer: u32,
1379    kda: &memra_gguf::model_plan::KimiDeltaNetPlan,
1380    hidden: usize,
1381) -> Result<(), ReferenceError> {
1382    let heads = kda.num_heads as usize;
1383    let head_dim = kda.head_dim as usize;
1384    let kernel = kda.conv_kernel as usize;
1385    let qkv = heads * head_dim;
1386    for (tensor, output, input, salt) in [
1387        (LayerTensor::KdaQuery, qkv, hidden, 140),
1388        (LayerTensor::KdaKey, qkv, hidden, 141),
1389        (LayerTensor::KdaValue, qkv, hidden, 142),
1390        (LayerTensor::KdaForgetDown, head_dim, hidden, 143),
1391        (LayerTensor::KdaForgetUp, qkv, head_dim, 144),
1392        (LayerTensor::KdaGateDown, head_dim, hidden, 145),
1393        (LayerTensor::KdaGateUp, qkv, head_dim, 146),
1394        (LayerTensor::KdaBeta, heads, hidden, 147),
1395        (LayerTensor::KdaOutput, hidden, qkv, 148),
1396    ] {
1397        weights.insert(
1398            layer_id(layer, tensor),
1399            generated_tensor(
1400                &[output, input],
1401                salt + layer as u64 * 163,
1402                1.0 / (input as f32).sqrt(),
1403            )?,
1404        );
1405    }
1406    for (tensor, salt) in [
1407        (LayerTensor::KdaQueryConv, 149),
1408        (LayerTensor::KdaKeyConv, 150),
1409        (LayerTensor::KdaValueConv, 151),
1410    ] {
1411        weights.insert(
1412            layer_id(layer, tensor),
1413            generated_tensor(
1414                &[qkv, kernel],
1415                salt + layer as u64 * 163,
1416                1.0 / (kernel as f32).sqrt(),
1417            )?,
1418        );
1419    }
1420    weights.insert(
1421        layer_id(layer, LayerTensor::KdaALog),
1422        generated_tensor(&[heads], 152 + layer as u64 * 163, 0.1)?,
1423    );
1424    // dt_bias is per-CHANNEL (width qkv), unlike GDN's per-head bias.
1425    weights.insert(
1426        layer_id(layer, LayerTensor::KdaDtBias),
1427        generated_tensor(&[qkv], 153 + layer as u64 * 163, 0.1)?,
1428    );
1429    weights.insert(
1430        layer_id(layer, LayerTensor::KdaOutputNorm),
1431        ReferenceTensor::new(vec![head_dim], vec![1.0; head_dim])?,
1432    );
1433    Ok(())
1434}
1435
1436fn add_mla_fixture(
1437    weights: &mut ReferenceWeights,
1438    layer: u32,
1439    mla: &memra_gguf::model_plan::MlaAttentionPlan,
1440    hidden: usize,
1441) -> Result<(), ReferenceError> {
1442    if let memra_gguf::model_plan::MlaAttentionPlan::CompressedKv { .. } = mla {
1443        return add_compressed_mla_fixture(weights, layer, mla, hidden);
1444    }
1445    let memra_gguf::model_plan::MlaAttentionPlan::LatentKv {
1446        query_heads,
1447        q_lora_rank,
1448        kv_lora_rank,
1449        qk_head_dim,
1450        rope_head_dim,
1451        value_head_dim,
1452        sparse_index,
1453        ..
1454    } = mla.clone()
1455    else {
1456        return Err(ReferenceError::UnsupportedOperation {
1457            layer: Some(layer),
1458            operation: "compressed-KV MLA fixture",
1459        });
1460    };
1461    let heads = query_heads as usize;
1462    let q_rank = q_lora_rank as usize;
1463    let kv_rank = kv_lora_rank as usize;
1464    let qk_dim = qk_head_dim as usize;
1465    let rope_dim = rope_head_dim as usize;
1466    let nope_dim = qk_dim - rope_dim;
1467    let value_dim = value_head_dim as usize;
1468    for (tensor, shape, input, salt) in [
1469        (LayerTensor::MlaQueryDown, vec![q_rank, hidden], hidden, 80),
1470        (
1471            LayerTensor::MlaQueryUp,
1472            vec![heads * qk_dim, q_rank],
1473            q_rank,
1474            81,
1475        ),
1476        (
1477            LayerTensor::MlaKvDown,
1478            vec![kv_rank + rope_dim, hidden],
1479            hidden,
1480            82,
1481        ),
1482        // `[head][rank][nope]`, NOT the checkpoint's own `[head][nope][rank]`. This is the
1483        // `TensorId::MlaKeyUp` layout the tensor contract declares (GGUF ne
1484        // `[nope, kv_rank, heads]`, fastest axis first), which is what llama.cpp's `attn_k_b`
1485        // mint and `hf_mapping::split_mla_kv_plane` both emit and what the engine's absorb GEMM
1486        // reads. One TensorId, one byte order: a fixture minted the other way round mis-strides
1487        // every engine-vs-reference MLA comparison while preserving element counts.
1488        (
1489            LayerTensor::MlaKeyUp,
1490            vec![heads, kv_rank, nope_dim],
1491            kv_rank,
1492            83,
1493        ),
1494        (
1495            LayerTensor::MlaValueUp,
1496            vec![heads, value_dim, kv_rank],
1497            kv_rank,
1498            84,
1499        ),
1500        (
1501            LayerTensor::MlaOutput,
1502            vec![hidden, heads * value_dim],
1503            heads * value_dim,
1504            85,
1505        ),
1506    ] {
1507        weights.insert(
1508            layer_id(layer, tensor),
1509            generated_tensor(
1510                &shape,
1511                salt + layer as u64 * 71,
1512                1.0 / (input as f32).sqrt(),
1513            )?,
1514        );
1515    }
1516    weights.insert(
1517        layer_id(layer, LayerTensor::MlaQueryDownNorm),
1518        ReferenceTensor::new(vec![q_rank], vec![1.0; q_rank])?,
1519    );
1520    weights.insert(
1521        layer_id(layer, LayerTensor::MlaKvDownNorm),
1522        ReferenceTensor::new(vec![kv_rank], vec![1.0; kv_rank])?,
1523    );
1524    // Only the k-pool indexer (glm5_next) owns tensors on the LatentKv path; the
1525    // per-token variant executes through full-selection equivalence without them.
1526    if let memra_gguf::model_plan::SparseIndexPlan::Own {
1527        heads: index_heads,
1528        head_dim: index_dim,
1529        top_k: _,
1530        kpool: Some(kpool),
1531    } = sparse_index
1532    {
1533        let index_heads = index_heads as usize;
1534        let index_dim = index_dim as usize;
1535        let pool = kpool.pool as usize;
1536        for (tensor, shape, input, salt) in [
1537            (
1538                LayerTensor::SparseQuery,
1539                vec![index_heads * index_dim, q_rank],
1540                q_rank,
1541                120,
1542            ),
1543            (LayerTensor::SparseKey, vec![index_dim, hidden], hidden, 121),
1544            (
1545                LayerTensor::SparseProjection,
1546                vec![index_heads, hidden],
1547                hidden,
1548                122,
1549            ),
1550            (
1551                LayerTensor::SparseCompressorGate,
1552                vec![index_dim, hidden],
1553                hidden,
1554                123,
1555            ),
1556            (
1557                LayerTensor::SparseCompressorPosition,
1558                vec![pool, index_dim],
1559                index_dim,
1560                124,
1561            ),
1562        ] {
1563            weights.insert(
1564                layer_id(layer, tensor),
1565                generated_tensor(
1566                    &shape,
1567                    salt + layer as u64 * 71,
1568                    1.0 / (input as f32).sqrt(),
1569                )?,
1570            );
1571        }
1572        weights.insert(
1573            layer_id(layer, LayerTensor::SparseKeyNorm),
1574            ReferenceTensor::new(vec![index_dim], vec![1.0; index_dim])?,
1575        );
1576        // LayerNorm bias (nonzero so the bias path is exercised).
1577        weights.insert(
1578            layer_id(layer, LayerTensor::SparseKeyNormBias),
1579            generated_tensor(&[index_dim], 125 + layer as u64 * 71, 0.05)?,
1580        );
1581    }
1582    Ok(())
1583}
1584
1585#[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
1586fn add_compressed_mla_fixture(
1587    weights: &mut ReferenceWeights,
1588    layer: u32,
1589    mla: &memra_gguf::model_plan::MlaAttentionPlan,
1590    hidden: usize,
1591) -> Result<(), ReferenceError> {
1592    use memra_gguf::model_plan::{MlaAttentionPlan, SparseIndexPlan};
1593
1594    let MlaAttentionPlan::CompressedKv {
1595        query_heads,
1596        q_lora_rank,
1597        latent_head_dim,
1598        rope_head_dim,
1599        output_lora_rank,
1600        output_groups,
1601        compressor,
1602        sparse_index,
1603        ..
1604    } = mla
1605    else {
1606        unreachable!()
1607    };
1608    let heads = *query_heads as usize;
1609    let q_rank = *q_lora_rank as usize;
1610    let head_dim = *latent_head_dim as usize;
1611    let rope_dim = *rope_head_dim as usize;
1612    let output_rank = *output_lora_rank as usize;
1613    let groups = *output_groups as usize;
1614    if groups == 0 || heads % groups != 0 || rope_dim > head_dim {
1615        return Err(ReferenceError::InvalidPlan {
1616            layer: Some(layer),
1617            reason: "compressed attention has invalid head or output-group geometry",
1618        });
1619    }
1620    let group_width = heads / groups * head_dim;
1621    for (tensor, shape, input, salt) in [
1622        (LayerTensor::MlaQueryDown, vec![q_rank, hidden], hidden, 110),
1623        (
1624            LayerTensor::MlaQueryUp,
1625            vec![heads * head_dim, q_rank],
1626            q_rank,
1627            111,
1628        ),
1629        (LayerTensor::MlaKvDown, vec![head_dim, hidden], hidden, 112),
1630        (
1631            LayerTensor::MlaOutputDown,
1632            vec![groups * output_rank, group_width],
1633            group_width,
1634            113,
1635        ),
1636        (
1637            LayerTensor::MlaOutput,
1638            vec![hidden, groups * output_rank],
1639            groups * output_rank,
1640            114,
1641        ),
1642    ] {
1643        weights.insert(
1644            layer_id(layer, tensor),
1645            generated_tensor(
1646                &shape,
1647                salt + layer as u64 * 131,
1648                1.0 / (input as f32).sqrt(),
1649            )?,
1650        );
1651    }
1652    weights.insert(
1653        layer_id(layer, LayerTensor::MlaQueryDownNorm),
1654        ReferenceTensor::new(vec![q_rank], vec![1.0; q_rank])?,
1655    );
1656    weights.insert(
1657        layer_id(layer, LayerTensor::MlaKvDownNorm),
1658        ReferenceTensor::new(vec![head_dim], vec![1.0; head_dim])?,
1659    );
1660    weights.insert(
1661        layer_id(layer, LayerTensor::AttentionSink),
1662        generated_tensor(&[heads], 115 + layer as u64 * 131, 0.05)?,
1663    );
1664    if let Some(compressor) = compressor {
1665        add_compressor_fixture(
1666            weights,
1667            layer,
1668            hidden,
1669            head_dim,
1670            compressor.ratio as usize,
1671            compressor.latent_dim as usize,
1672            false,
1673        )?;
1674    }
1675    match sparse_index {
1676        SparseIndexPlan::None => {}
1677        SparseIndexPlan::Own {
1678            heads, head_dim, ..
1679        } => {
1680            let Some(compressor) = compressor else {
1681                return Err(ReferenceError::InvalidPlan {
1682                    layer: Some(layer),
1683                    reason: "compressed sparse index requires a compressor ratio",
1684                });
1685            };
1686            let index_heads = *heads as usize;
1687            let index_dim = *head_dim as usize;
1688            weights.insert(
1689                layer_id(layer, LayerTensor::SparseQuery),
1690                generated_tensor(
1691                    &[index_heads * index_dim, q_rank],
1692                    116 + layer as u64 * 131,
1693                    1.0 / (q_rank as f32).sqrt(),
1694                )?,
1695            );
1696            weights.insert(
1697                layer_id(layer, LayerTensor::SparseProjection),
1698                generated_tensor(
1699                    &[index_heads, hidden],
1700                    117 + layer as u64 * 131,
1701                    1.0 / (hidden as f32).sqrt(),
1702                )?,
1703            );
1704            add_compressor_fixture(
1705                weights,
1706                layer,
1707                hidden,
1708                index_dim,
1709                compressor.ratio as usize,
1710                2 * index_dim,
1711                true,
1712            )?;
1713        }
1714        SparseIndexPlan::SharedFromPrevious { .. } => {
1715            return Err(ReferenceError::UnsupportedOperation {
1716                layer: Some(layer),
1717                operation: "shared compressed sparse-index fixture",
1718            });
1719        }
1720    }
1721    Ok(())
1722}
1723
1724#[allow(clippy::too_many_arguments)]
1725fn add_compressor_fixture(
1726    weights: &mut ReferenceWeights,
1727    layer: u32,
1728    hidden: usize,
1729    output_dim: usize,
1730    ratio: usize,
1731    latent: usize,
1732    sparse: bool,
1733) -> Result<(), ReferenceError> {
1734    let (key_value, gate, norm, position, salt) = if sparse {
1735        (
1736            LayerTensor::SparseCompressorKeyValue,
1737            LayerTensor::SparseCompressorGate,
1738            LayerTensor::SparseCompressorNorm,
1739            LayerTensor::SparseCompressorPosition,
1740            121,
1741        )
1742    } else {
1743        (
1744            LayerTensor::KvCompressorKeyValue,
1745            LayerTensor::KvCompressorGate,
1746            LayerTensor::KvCompressorNorm,
1747            LayerTensor::KvCompressorPosition,
1748            118,
1749        )
1750    };
1751    for (tensor, offset) in [(key_value, 0), (gate, 1)] {
1752        weights.insert(
1753            layer_id(layer, tensor),
1754            generated_tensor(
1755                &[latent, hidden],
1756                salt + offset + layer as u64 * 131,
1757                1.0 / (hidden as f32).sqrt(),
1758            )?,
1759        );
1760    }
1761    weights.insert(
1762        layer_id(layer, norm),
1763        ReferenceTensor::new(vec![output_dim], vec![1.0; output_dim])?,
1764    );
1765    weights.insert(
1766        layer_id(layer, position),
1767        generated_tensor(&[ratio, latent], salt + 2 + layer as u64 * 131, 0.05)?,
1768    );
1769    Ok(())
1770}
1771
1772fn add_dense_mlp_fixture(
1773    weights: &mut ReferenceWeights,
1774    layer: u32,
1775    mlp: &memra_gguf::model_plan::DenseMlpPlan,
1776    hidden: usize,
1777) -> Result<(), ReferenceError> {
1778    let intermediate = mlp.intermediate_size as usize;
1779    for (tensor, output, input, salt) in [
1780        (LayerTensor::MlpGate, intermediate, hidden, 20),
1781        (LayerTensor::MlpUp, intermediate, hidden, 21),
1782        (LayerTensor::MlpDown, hidden, intermediate, 22),
1783    ] {
1784        weights.insert(
1785            layer_id(layer, tensor),
1786            generated_tensor(
1787                &[output, input],
1788                salt + layer as u64 * 31,
1789                1.0 / (input as f32).sqrt(),
1790            )?,
1791        );
1792    }
1793    Ok(())
1794}
1795
1796fn add_moe_fixture(
1797    weights: &mut ReferenceWeights,
1798    layer: u32,
1799    moe: &memra_gguf::model_plan::MoeMlpPlan,
1800    hidden: usize,
1801    vocab: usize,
1802) -> Result<(), ReferenceError> {
1803    let experts = moe.expert_count as usize;
1804    let selected = moe.experts_per_token as usize;
1805    let intermediate = moe.expert_intermediate_size as usize;
1806    if matches!(
1807        moe.router,
1808        memra_gguf::model_plan::RouterPlan::TokenIdHash { .. }
1809    ) {
1810        let mut table = Vec::with_capacity(vocab * selected);
1811        for token in 0..vocab {
1812            for rank in 0..selected {
1813                table.push(((token + rank) % experts) as f32);
1814            }
1815        }
1816        weights.insert(
1817            layer_id(layer, LayerTensor::MoeTokenToExpert),
1818            ReferenceTensor::new(vec![vocab, selected], table)?,
1819        );
1820    }
1821    weights.insert(
1822        layer_id(layer, LayerTensor::MoeRouter),
1823        generated_tensor(
1824            &[experts, hidden],
1825            60 + layer as u64 * 59,
1826            1.0 / (hidden as f32).sqrt(),
1827        )?,
1828    );
1829    if router_has_selection_bias(&moe.router) {
1830        weights.insert(
1831            layer_id(layer, LayerTensor::MoeRouterBias),
1832            generated_tensor(&[experts], 61 + layer as u64 * 59, 0.05)?,
1833        );
1834    }
1835    for (tensor, shape, input, salt) in [
1836        (
1837            LayerTensor::MoeExpertGateBank,
1838            vec![experts, intermediate, hidden],
1839            hidden,
1840            62,
1841        ),
1842        (
1843            LayerTensor::MoeExpertUpBank,
1844            vec![experts, intermediate, hidden],
1845            hidden,
1846            63,
1847        ),
1848        (
1849            LayerTensor::MoeExpertDownBank,
1850            vec![experts, hidden, intermediate],
1851            intermediate,
1852            64,
1853        ),
1854    ] {
1855        weights.insert(
1856            layer_id(layer, tensor),
1857            generated_tensor(
1858                &shape,
1859                salt + layer as u64 * 59,
1860                1.0 / (input as f32).sqrt(),
1861            )?,
1862        );
1863    }
1864    if let Some(shared) = moe.shared.as_ref() {
1865        let intermediate = shared.intermediate_size as usize;
1866        for (tensor, output, input, salt) in [
1867            (LayerTensor::SharedMlpGate, intermediate, hidden, 65),
1868            (LayerTensor::SharedMlpUp, intermediate, hidden, 66),
1869            (LayerTensor::SharedMlpDown, hidden, intermediate, 67),
1870        ] {
1871            weights.insert(
1872                layer_id(layer, tensor),
1873                generated_tensor(
1874                    &[output, input],
1875                    salt + layer as u64 * 59,
1876                    1.0 / (input as f32).sqrt(),
1877                )?,
1878            );
1879        }
1880        if shared.gated {
1881            weights.insert(
1882                layer_id(layer, LayerTensor::SharedMlpInputGate),
1883                generated_tensor(&[hidden], 68 + layer as u64 * 59, 0.2)?,
1884            );
1885        }
1886    }
1887    Ok(())
1888}
1889
1890fn add_gemma_parallel_moe_fixture(
1891    weights: &mut ReferenceWeights,
1892    layer: u32,
1893    moe: &memra_gguf::model_plan::MoeMlpPlan,
1894    hidden: usize,
1895) -> Result<(), ReferenceError> {
1896    let experts = moe.expert_count as usize;
1897    let intermediate = moe.expert_intermediate_size as usize;
1898    weights.insert(
1899        layer_id(layer, LayerTensor::MoeExpertGateUpBank),
1900        generated_tensor(
1901            &[experts, 2 * intermediate, hidden],
1902            180 + layer as u64 * 19,
1903            1.0 / (hidden as f32).sqrt(),
1904        )?,
1905    );
1906    weights.insert(
1907        layer_id(layer, LayerTensor::MoeRouterScale),
1908        ReferenceTensor::new(vec![hidden], vec![1.0; hidden])?,
1909    );
1910    weights.insert(
1911        layer_id(layer, LayerTensor::MoeExpertOutputScale),
1912        generated_tensor(&[experts], 181 + layer as u64 * 19, 0.2)?,
1913    );
1914    Ok(())
1915}
1916
1917pub fn execute(
1918    plan: &ModelPlan,
1919    weights: &ReferenceWeights,
1920    token_ids: &[u32],
1921) -> Result<ReferenceOutput, ReferenceError> {
1922    if token_ids.is_empty() {
1923        return Err(ReferenceError::EmptyInput);
1924    }
1925    let hidden = plan.hidden_size as usize;
1926    let vocab = plan.vocab_size as usize;
1927    let embedding = tensor(weights, &TensorId::TokenEmbedding, &[vocab, hidden])?;
1928    let embedded = embed_token_rows(plan, embedding, token_ids, vocab, hidden)?;
1929    execute_embedded(plan, weights, token_ids, embedding, embedded)
1930}
1931
1932pub fn execute_multimodal(
1933    plan: &ModelPlan,
1934    weights: &ReferenceWeights,
1935    token_ids: &[u32],
1936    vision_input: &ReferenceVisionInput,
1937) -> Result<ReferenceMultimodalOutput, ReferenceError> {
1938    if token_ids.is_empty() {
1939        return Err(ReferenceError::EmptyInput);
1940    }
1941    let injection = plan.multimodal.ok_or(ReferenceError::InvalidPlan {
1942        layer: None,
1943        reason: "multimodal input requires a vision-token injection plan",
1944    })?;
1945    let vision = execute_vision(plan, weights, vision_input)?;
1946    if let Some(tokens_per_image) = injection.tokens_per_image
1947        && vision.output_tokens != tokens_per_image as usize
1948    {
1949        return Err(ReferenceError::InvalidPlan {
1950            layer: None,
1951            reason: "vision output token count does not match the injection plan",
1952        });
1953    }
1954    let placeholder_count = token_ids
1955        .iter()
1956        .filter(|&&token| token == injection.placeholder_token_id)
1957        .count();
1958    if placeholder_count != vision.output_tokens {
1959        return Err(ReferenceError::InvalidPlan {
1960            layer: None,
1961            reason: "image placeholder count does not match projected vision tokens",
1962        });
1963    }
1964    let hidden = plan.hidden_size as usize;
1965    let vocab = plan.vocab_size as usize;
1966    let embedding = tensor(weights, &TensorId::TokenEmbedding, &[vocab, hidden])?;
1967    let mut embedded = embed_token_rows(plan, embedding, token_ids, vocab, hidden)?;
1968    let mut vision_row = 0;
1969    for (position, &token) in token_ids.iter().enumerate() {
1970        if token == injection.placeholder_token_id {
1971            embedded[position * hidden..(position + 1) * hidden].copy_from_slice(
1972                &vision.projected_hidden[vision_row * hidden..(vision_row + 1) * hidden],
1973            );
1974            vision_row += 1;
1975        }
1976    }
1977    let language = execute_embedded(plan, weights, token_ids, embedding, embedded)?;
1978    Ok(ReferenceMultimodalOutput { language, vision })
1979}
1980
1981fn embed_token_rows(
1982    plan: &ModelPlan,
1983    embedding: &[f32],
1984    token_ids: &[u32],
1985    vocab: usize,
1986    hidden: usize,
1987) -> Result<Vec<f32>, ReferenceError> {
1988    let mut embedded = vec![0.0; token_ids.len() * hidden];
1989    for (position, &token) in token_ids.iter().enumerate() {
1990        let token = token as usize;
1991        if token >= vocab {
1992            return Err(ReferenceError::TokenOutOfRange {
1993                token: token as u32,
1994                vocab,
1995            });
1996        }
1997        embedded[position * hidden..(position + 1) * hidden]
1998            .copy_from_slice(&embedding[token * hidden..(token + 1) * hidden]);
1999        if plan.embedding_scale != 1.0 {
2000            for value in &mut embedded[position * hidden..(position + 1) * hidden] {
2001                *value *= plan.embedding_scale;
2002            }
2003        }
2004    }
2005    Ok(embedded)
2006}
2007
2008fn execute_embedded(
2009    plan: &ModelPlan,
2010    weights: &ReferenceWeights,
2011    token_ids: &[u32],
2012    embedding: &[f32],
2013    embedded: Vec<f32>,
2014) -> Result<ReferenceOutput, ReferenceError> {
2015    let tokens = token_ids.len();
2016    let hidden = plan.hidden_size as usize;
2017    let vocab = plan.vocab_size as usize;
2018    if embedded.len() != tokens * hidden {
2019        return Err(ReferenceError::InvalidPlan {
2020            layer: None,
2021            reason: "embedded language input does not match tokens x hidden",
2022        });
2023    }
2024    let hyper = hyper_topology(plan)?;
2025    let gated = gated_residual_topology(plan)?;
2026    let mut x = if let Some((streams, _)) = gated {
2027        // qwen4_exp entry: the wide stream starts as `streams` copies of the embedding
2028        // (modular L1012 `repeat(1, 1, hc_count)`).
2029        let wide = streams * hidden;
2030        let mut expanded = vec![0.0; tokens * wide];
2031        for token in 0..tokens {
2032            for stream in 0..streams {
2033                expanded[token * wide + stream * hidden..token * wide + (stream + 1) * hidden]
2034                    .copy_from_slice(&embedded[token * hidden..(token + 1) * hidden]);
2035            }
2036        }
2037        expanded
2038    } else if let Some((streams, _, _, _)) = hyper {
2039        memra_gguf::dsv4_forward::hc_expand(&embedded, tokens, streams, hidden)
2040    } else {
2041        embedded.clone()
2042    };
2043
2044    let mut state = Vec::with_capacity(plan.layers.len());
2045    let mut layer_hidden = Vec::with_capacity(plan.layers.len());
2046    let dspark = plan.drafter.as_ref().map(|drafter| match drafter {
2047        memra_gguf::model_plan::DrafterPlan::Dspark(plan) => plan,
2048    });
2049    let mut draft_taps = dspark.map(|plan| vec![None; plan.target_layer_ids.len()]);
2050    for layer in &plan.layers {
2051        let (next, layer_state) = execute_layer(
2052            layer,
2053            weights,
2054            &x,
2055            token_ids,
2056            tokens,
2057            hidden,
2058            vocab,
2059            LayerScope::Trunk,
2060        )?;
2061        x = next;
2062        layer_hidden.push(x.clone());
2063        if let (Some(dspark), Some(taps)) = (dspark, draft_taps.as_mut())
2064            && let Some(target) = dspark
2065                .target_layer_ids
2066                .iter()
2067                .position(|&target| target == layer.index)
2068        {
2069            taps[target] = Some(collapse_stream_mean(
2070                &x,
2071                tokens,
2072                hidden,
2073                hyper.map(|topology| topology.0),
2074            )?);
2075        }
2076        state.push(layer_state);
2077    }
2078    let trunk_hidden = x.clone();
2079    let output = weights
2080        .get(&TensorId::OutputProjection)
2081        .map(|tensor| tensor_checked(&TensorId::OutputProjection, tensor, &[vocab, hidden]))
2082        .transpose()?
2083        .unwrap_or(embedding);
2084    let logits = if let Some((streams, rank)) = gated {
2085        // Exit downmix replaces the final norm: the mixer read gate (use_combine=false)
2086        // collapses the wide stream and its grouped hc_norm IS the exit normalization
2087        // (SEMANTICS.md §Layer stack; census has no model.language_model.norm), so this
2088        // arm bypasses project_trunk_logits' OutputNorm rms_norm.
2089        let collapsed = gated_residual_read(
2090            weights,
2091            LayerScope::Trunk.mixer_prefix(),
2092            "",
2093            &trunk_hidden,
2094            tokens,
2095            streams,
2096            hidden,
2097            rank,
2098            plan.output_norm.epsilon,
2099            false,
2100        )?
2101        .0;
2102        let mut logits = linear(&collapsed, output, tokens, hidden, vocab);
2103        apply_logits_transforms(&mut logits, vocab, &plan.logits);
2104        logits
2105    } else {
2106        project_trunk_logits(
2107            plan,
2108            weights,
2109            &trunk_hidden,
2110            tokens,
2111            hidden,
2112            vocab,
2113            embedding,
2114        )?
2115    };
2116    let draft = match (dspark, draft_taps) {
2117        (Some(dspark), Some(taps)) => Some(execute_dspark(
2118            dspark,
2119            weights,
2120            token_ids,
2121            embedding,
2122            output,
2123            &plan.logits,
2124            plan.output_norm.epsilon,
2125            hidden,
2126            vocab,
2127            taps,
2128        )?),
2129        _ => None,
2130    };
2131    // MTP fusion consumes the COLLAPSED pre-output_norm hidden — the same collapse the
2132    // LM-head projection above just applied and the same quantity the engine hands over as
2133    // `h_seed` (MTP-PLAN §A). Passing the raw `[tokens, streams*hidden]` stream stack was
2134    // the reference-side refusal that kept `execute_mtp` erroring on every hc plan
2135    // ("HyperConnections MTP fusion") while the plan, the contract, and the checkpoint all
2136    // carried the NextN block.
2137    let mtp_hidden = collapse_trunk_hidden(plan, weights, &trunk_hidden, tokens, hidden)?;
2138    let mtp = execute_mtp(
2139        plan,
2140        weights,
2141        token_ids,
2142        embedding,
2143        mtp_hidden.as_deref().unwrap_or(&trunk_hidden),
2144        tokens,
2145        hidden,
2146        vocab,
2147        output,
2148    )?;
2149    Ok(ReferenceOutput {
2150        logits,
2151        tokens,
2152        vocab,
2153        state: ReferenceState { layers: state },
2154        mtp,
2155        draft,
2156        layer_hidden,
2157    })
2158}
2159
2160/// Trunk exit shared by [`execute`] and the streamed driver: stream collapse, output
2161/// norm, LM-head projection, and logits transforms. `embedding` is the tied-head
2162/// fallback when `OutputProjection` is absent. Kept as one function so the two paths
2163/// cannot drift; the checkpoint runner's `--self-test` pins them bit-for-bit.
2164/// The trunk-exit stream collapse: `[tokens, streams*hidden]` -> `[tokens, hidden]` for an
2165/// hc plan, identity for a serial one. ONE function for both consumers — the LM-head
2166/// projection and the MTP fusion input — because the engine's `h_seed` contract (MTP-PLAN
2167/// §A) is "the PRE-output_norm hidden, taken from the collapsed stack so it means the same
2168/// thing it does on the serial path": if the two collapses could drift, the MTP oracle
2169/// would gate the draft against a hidden the trunk never hands over.
2170fn collapse_trunk_hidden(
2171    plan: &ModelPlan,
2172    weights: &ReferenceWeights,
2173    trunk_hidden: &[f32],
2174    tokens: usize,
2175    hidden: usize,
2176) -> Result<Option<Vec<f32>>, ReferenceError> {
2177    let Some((streams, epsilon, _, collapse)) = hyper_topology(plan)? else {
2178        return Ok(None);
2179    };
2180    Ok(Some(match collapse {
2181        HcCollapse::GatedHead => collapse_hyper_head(
2182            weights,
2183            trunk_hidden,
2184            tokens,
2185            streams,
2186            hidden,
2187            plan,
2188            epsilon,
2189        )?,
2190        HcCollapse::Mean => collapse_stream_mean(trunk_hidden, tokens, hidden, Some(streams))?,
2191    }))
2192}
2193
2194fn project_trunk_logits(
2195    plan: &ModelPlan,
2196    weights: &ReferenceWeights,
2197    trunk_hidden: &[f32],
2198    tokens: usize,
2199    hidden: usize,
2200    vocab: usize,
2201    embedding: &[f32],
2202) -> Result<Vec<f32>, ReferenceError> {
2203    let collapsed = collapse_trunk_hidden(plan, weights, trunk_hidden, tokens, hidden)?;
2204    let x: &[f32] = collapsed.as_deref().unwrap_or(trunk_hidden);
2205    if crate::hidden_trace::enabled() {
2206        crate::hidden_trace::emit_last_row("collapse", -1, tokens, hidden, x);
2207    }
2208    let x = rms_norm(
2209        x,
2210        tokens,
2211        hidden,
2212        tensor(weights, &TensorId::OutputNorm, &[hidden])?,
2213        plan.output_norm.epsilon,
2214    );
2215    let output = weights
2216        .get(&TensorId::OutputProjection)
2217        .map(|tensor| tensor_checked(&TensorId::OutputProjection, tensor, &[vocab, hidden]))
2218        .transpose()?
2219        .unwrap_or(embedding);
2220    let mut logits = linear(&x, output, tokens, hidden, vocab);
2221    apply_logits_transforms(&mut logits, vocab, &plan.logits);
2222    Ok(logits)
2223}
2224
2225/// Layer-at-a-time trunk execution over the exact per-layer math of [`execute`], for
2226/// checkpoint-scale runs where all weights cannot be resident at once. The driver
2227/// materializes only the current layer's tensors, calls [`StreamedTrunkExecution::step`],
2228/// and frees them before the next layer.
2229///
2230/// `begin` needs `TokenEmbedding` in `globals`; `finish` needs `OutputNorm` plus
2231/// `OutputProjection` (falling back to the embedding for tied heads) and, for a
2232/// gated-head collapse, the `HyperHead*` tensors.
2233///
2234/// Deliberate scope: trunk + final norm + LM head only. MTP blocks are not executed
2235/// (`mtp` stays empty) and drafter plans are refused at `begin` — the streamed path has
2236/// no per-layer tap capture. The glm5 checkpoint runner's `--self-test` mode pins this
2237/// path against [`execute`] bit-for-bit.
2238pub struct StreamedTrunkExecution<'a> {
2239    plan: &'a ModelPlan,
2240    token_ids: Vec<u32>,
2241    x: Vec<f32>,
2242    tokens: usize,
2243    hidden: usize,
2244    vocab: usize,
2245    next: usize,
2246    states: Vec<ReferenceLayerState>,
2247}
2248
2249impl<'a> StreamedTrunkExecution<'a> {
2250    pub fn begin(
2251        plan: &'a ModelPlan,
2252        globals: &ReferenceWeights,
2253        token_ids: &[u32],
2254    ) -> Result<Self, ReferenceError> {
2255        if token_ids.is_empty() {
2256            return Err(ReferenceError::EmptyInput);
2257        }
2258        if plan.drafter.is_some() {
2259            return Err(ReferenceError::UnsupportedOperation {
2260                layer: None,
2261                operation: "streamed drafter execution",
2262            });
2263        }
2264        let tokens = token_ids.len();
2265        let hidden = plan.hidden_size as usize;
2266        let vocab = plan.vocab_size as usize;
2267        let embedding = tensor(globals, &TensorId::TokenEmbedding, &[vocab, hidden])?;
2268        let embedded = embed_token_rows(plan, embedding, token_ids, vocab, hidden)?;
2269        let x = match hyper_topology(plan)? {
2270            Some((streams, _, _, _)) => {
2271                memra_gguf::dsv4_forward::hc_expand(&embedded, tokens, streams, hidden)
2272            }
2273            None => embedded,
2274        };
2275        if crate::hidden_trace::enabled() {
2276            crate::hidden_trace::emit_tokens(token_ids);
2277            let width = x.len() / tokens;
2278            crate::hidden_trace::emit_last_row("expand", -1, tokens, width, &x);
2279        }
2280        Ok(Self {
2281            plan,
2282            token_ids: token_ids.to_vec(),
2283            x,
2284            tokens,
2285            hidden,
2286            vocab,
2287            next: 0,
2288            states: Vec::with_capacity(plan.layers.len()),
2289        })
2290    }
2291
2292    /// The plan layer the next [`Self::step`] call will execute, or `None` when the
2293    /// trunk is fully executed.
2294    pub fn next_layer(&self) -> Option<&'a memra_gguf::model_plan::LayerPlan> {
2295        self.plan.layers.get(self.next)
2296    }
2297
2298    /// Execute the next trunk layer using only that layer's tensors. Returns the
2299    /// executed layer's plan index.
2300    pub fn step(&mut self, weights: &ReferenceWeights) -> Result<u32, ReferenceError> {
2301        let layer = self
2302            .plan
2303            .layers
2304            .get(self.next)
2305            .ok_or(ReferenceError::InvalidPlan {
2306                layer: None,
2307                reason: "streamed trunk stepped past the final layer",
2308            })?;
2309        let (next, layer_state) = execute_layer(
2310            layer,
2311            weights,
2312            &self.x,
2313            &self.token_ids,
2314            self.tokens,
2315            self.hidden,
2316            self.vocab,
2317            LayerScope::Trunk,
2318        )?;
2319        self.x = next;
2320        self.states.push(layer_state);
2321        self.next += 1;
2322        Ok(layer.index)
2323    }
2324
2325    /// Collapse, final-norm, and project the trunk. MTP blocks are skipped by design.
2326    pub fn finish(self, globals: &ReferenceWeights) -> Result<ReferenceOutput, ReferenceError> {
2327        if self.next != self.plan.layers.len() {
2328            return Err(ReferenceError::InvalidPlan {
2329                layer: None,
2330                reason: "streamed trunk finished before executing every layer",
2331            });
2332        }
2333        let embedding = tensor(
2334            globals,
2335            &TensorId::TokenEmbedding,
2336            &[self.vocab, self.hidden],
2337        )?;
2338        let logits = project_trunk_logits(
2339            self.plan,
2340            globals,
2341            &self.x,
2342            self.tokens,
2343            self.hidden,
2344            self.vocab,
2345            embedding,
2346        )?;
2347        Ok(ReferenceOutput {
2348            logits,
2349            tokens: self.tokens,
2350            vocab: self.vocab,
2351            state: ReferenceState {
2352                layers: self.states,
2353            },
2354            mtp: Vec::new(),
2355            draft: None,
2356            // Streamed runs are checkpoint-scale by definition; retaining every
2357            // layer's residual would defeat the memory bound. Parity localization
2358            // uses the in-memory [`execute`] path.
2359            layer_hidden: Vec::new(),
2360        })
2361    }
2362}
2363
2364pub fn execute_vision(
2365    plan: &ModelPlan,
2366    weights: &ReferenceWeights,
2367    input: &ReferenceVisionInput,
2368) -> Result<ReferenceVisionOutput, ReferenceError> {
2369    let Some(vision) = plan.vision.as_ref() else {
2370        return Err(ReferenceError::InvalidPlan {
2371            layer: None,
2372            reason: "vision input requires a vision subplan",
2373        });
2374    };
2375    let vision = match vision {
2376        memra_gguf::model_plan::VisionPlan::Factored(vision) => vision,
2377        memra_gguf::model_plan::VisionPlan::Glm5Fused(vision) => {
2378            return execute_vision_glm5(vision, weights, input);
2379        }
2380    };
2381    if vision.clipped_linears {
2382        return Err(ReferenceError::UnsupportedOperation {
2383            layer: None,
2384            operation: "clipped vision linears",
2385        });
2386    }
2387    let patches = input.positions.len();
2388    let hidden = vision.hidden_size as usize;
2389    let patch_width =
2390        (vision.patch.channels * vision.patch.patch_size * vision.patch.patch_size) as usize;
2391    if input.patches.shape != [patches, patch_width]
2392        || input.output_tokens == 0
2393        || input.output_tokens > patches
2394    {
2395        return Err(ReferenceError::InvalidPlan {
2396            layer: None,
2397            reason: "vision patch input shape or output-token count is invalid",
2398        });
2399    }
2400    let mut normalized_patches = input.patches.data.clone();
2401    for value in &mut normalized_patches {
2402        *value = 2.0 * (*value - 0.5);
2403    }
2404    let mut x = linear(
2405        &normalized_patches,
2406        tensor(
2407            weights,
2408            &TensorId::Vision {
2409                layer: None,
2410                tensor: VisionTensor::PatchProjection,
2411            },
2412            &[hidden, patch_width],
2413        )?,
2414        patches,
2415        patch_width,
2416        hidden,
2417    );
2418    let position_table = tensor(
2419        weights,
2420        &TensorId::Vision {
2421            layer: None,
2422            tensor: VisionTensor::PositionEmbedding,
2423        },
2424        &[
2425            vision.patch.position_axes as usize,
2426            vision.patch.position_embedding_size as usize,
2427            hidden,
2428        ],
2429    )?;
2430    for (patch, position) in input.positions.iter().enumerate() {
2431        for (axis, &coordinate) in position.iter().enumerate() {
2432            let coordinate = coordinate as usize;
2433            if axis >= vision.patch.position_axes as usize
2434                || coordinate >= vision.patch.position_embedding_size as usize
2435            {
2436                return Err(ReferenceError::InvalidPlan {
2437                    layer: None,
2438                    reason: "vision patch position is outside the embedding table",
2439                });
2440            }
2441            let source =
2442                (axis * vision.patch.position_embedding_size as usize + coordinate) * hidden;
2443            for column in 0..hidden {
2444                x[patch * hidden + column] += position_table[source + column];
2445            }
2446        }
2447    }
2448    for layer in &vision.layers {
2449        x = execute_vision_layer(layer, weights, &x, &input.positions, patches, hidden)?;
2450    }
2451    let encoder_hidden = x.clone();
2452    let pooled_hidden = vision_pool(&x, &input.positions, patches, input.output_tokens, hidden)?;
2453    let mut standardized = pooled_hidden.clone();
2454    if vision.standardize {
2455        let bias = tensor(
2456            weights,
2457            &TensorId::Vision {
2458                layer: None,
2459                tensor: VisionTensor::StandardizeBias,
2460            },
2461            &[hidden],
2462        )?;
2463        let scale = tensor(
2464            weights,
2465            &TensorId::Vision {
2466                layer: None,
2467                tensor: VisionTensor::StandardizeScale,
2468            },
2469            &[hidden],
2470        )?;
2471        for row in standardized.chunks_exact_mut(hidden) {
2472            for column in 0..hidden {
2473                row[column] = (row[column] - bias[column]) * scale[column];
2474            }
2475        }
2476    }
2477    let standardized = rms_norm(
2478        &standardized,
2479        input.output_tokens,
2480        hidden,
2481        &vec![1.0; hidden],
2482        vision.layers[0].input_norm.epsilon,
2483    );
2484    let projection_size = vision.projection_output_size as usize;
2485    let projected_hidden = linear(
2486        &standardized,
2487        tensor(
2488            weights,
2489            &TensorId::Vision {
2490                layer: None,
2491                tensor: VisionTensor::OutputProjection,
2492            },
2493            &[projection_size, hidden],
2494        )?,
2495        input.output_tokens,
2496        hidden,
2497        projection_size,
2498    );
2499    Ok(ReferenceVisionOutput {
2500        encoder_hidden,
2501        pooled_hidden,
2502        projected_hidden,
2503        patch_count: patches,
2504        output_tokens: input.output_tokens,
2505        hidden_size: hidden,
2506        projection_size,
2507    })
2508}
2509
2510/// glm5_next tower forward. Semantics pinned against transformers 5.16.1
2511/// `Glm5NextVisionModel.forward` (vision classes diffed byte-identical to transformers
2512/// main; lane research/glm5-vision-20260830): patch conv as a linear over `(c, t, ph, pw)`
2513/// rows, per-head q/k RMS norms BEFORE the 2D rope, rope-only positions (theta 10000,
2514/// h-half then w-half, NeoX pairs `(d, d + head_dim/2)`), scaled non-causal attention,
2515/// biased clamped-SwiGLU block MLPs, post-encoder RMS norm, conv `merge x merge`
2516/// downsample over block-major token groups, gated clamped merger
2517/// (proj -> LayerNorm -> exact GELU -> clamp(gate,up) -> silu(gate)*up -> down).
2518#[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
2519fn execute_vision_glm5(
2520    vision: &memra_gguf::model_plan::Glm5VisionPlan,
2521    weights: &ReferenceWeights,
2522    input: &ReferenceVisionInput,
2523) -> Result<ReferenceVisionOutput, ReferenceError> {
2524    let hidden = vision.hidden_size as usize;
2525    let heads = vision.heads as usize;
2526    let head_dim = vision.head_dim as usize;
2527    let ff = vision.intermediate_size as usize;
2528    let out_width = vision.out_hidden_size as usize;
2529    let proj_inter = vision.projection_intermediate_size as usize;
2530    let merge = vision.spatial_merge_size as usize;
2531    let merge_area = merge * merge;
2532    let patch_width = vision.patch_input_width as usize;
2533    let limit = vision.swiglu_limit;
2534    let eps = vision.norm.epsilon;
2535    let tokens = input.positions.len();
2536    if input.patches.shape != [tokens, patch_width]
2537        || tokens == 0
2538        || tokens % merge_area != 0
2539        || input.output_tokens != tokens / merge_area
2540    {
2541        return Err(ReferenceError::InvalidPlan {
2542            layer: None,
2543            reason: "glm5 vision patch input shape, merge alignment or token count is invalid",
2544        });
2545    }
2546    let id = |layer: Option<u32>, tensor| TensorId::Vision { layer, tensor };
2547    let tensor_5d = |tensor, expected: &[usize]| -> Result<&[f32], ReferenceError> {
2548        tensor_checked(
2549            &id(None, tensor),
2550            weights
2551                .get(&id(None, tensor))
2552                .ok_or(ReferenceError::MissingTensor(id(None, tensor)))?,
2553            expected,
2554        )
2555    };
2556    // Patch embed: conv3d [hidden, c, t, ph, pw] row-major == linear rows over the
2557    // processor's (c, t, ph, pw) flat patch order.
2558    let patch_weight = tensor_5d(
2559        VisionTensor::PatchProjection,
2560        &[
2561            hidden,
2562            vision.in_channels as usize,
2563            vision.temporal_patch_size as usize,
2564            vision.patch_size as usize,
2565            vision.patch_size as usize,
2566        ],
2567    )?;
2568    let patch_bias = tensor(
2569        weights,
2570        &id(None, VisionTensor::PatchProjectionBias),
2571        &[hidden],
2572    )?;
2573    let mut x = linear(
2574        &input.patches.data,
2575        patch_weight,
2576        tokens,
2577        patch_width,
2578        hidden,
2579    );
2580    for row in x.chunks_exact_mut(hidden) {
2581        add_in_place(row, patch_bias);
2582    }
2583    // 2D rope tables: half rotates by h, half by w; inv_freq[i] = theta^(-2i/half),
2584    // NeoX pairs (d, d + half) share cos/sin (upstream cat((rotary, rotary), -1)).
2585    let half = head_dim / 2;
2586    let quarter = half / 2;
2587    let inv_freq: Vec<f32> = (0..quarter)
2588        .map(|index| vision.rope_theta.powf(-((2 * index) as f32) / half as f32))
2589        .collect();
2590    let mut rope_cos = vec![0.0f32; tokens * half];
2591    let mut rope_sin = vec![0.0f32; tokens * half];
2592    for (token, position) in input.positions.iter().enumerate() {
2593        for dim in 0..half {
2594            let angle = if dim < quarter {
2595                position[0] as f32 * inv_freq[dim]
2596            } else {
2597                position[1] as f32 * inv_freq[dim - quarter]
2598            };
2599            rope_cos[token * half + dim] = angle.cos();
2600            rope_sin[token * half + dim] = angle.sin();
2601        }
2602    }
2603    for layer in 0..vision.depth {
2604        let l = Some(layer);
2605        let layer_tensor =
2606            |tensor: VisionTensor, expected: &[usize]| -> Result<&[f32], ReferenceError> {
2607                self::tensor(weights, &id(l, tensor), expected)
2608            };
2609        // attn: rms(norm1) -> fused qkv+bias -> per-head q/k RMS -> rope -> sdpa -> proj+bias
2610        let attention_input = rms_norm(
2611            &x,
2612            tokens,
2613            hidden,
2614            layer_tensor(VisionTensor::InputNorm, &[hidden])?,
2615            eps,
2616        );
2617        let mut qkv = linear(
2618            &attention_input,
2619            layer_tensor(VisionTensor::FusedQkv, &[3 * hidden, hidden])?,
2620            tokens,
2621            hidden,
2622            3 * hidden,
2623        );
2624        let qkv_bias = layer_tensor(VisionTensor::FusedQkvBias, &[3 * hidden])?;
2625        for row in qkv.chunks_exact_mut(3 * hidden) {
2626            add_in_place(row, qkv_bias);
2627        }
2628        let query_norm = layer_tensor(VisionTensor::QueryNorm, &[head_dim])?;
2629        let key_norm = layer_tensor(VisionTensor::KeyNorm, &[head_dim])?;
2630        let mut query = vec![0.0f32; tokens * hidden];
2631        let mut key = vec![0.0f32; tokens * hidden];
2632        let mut value = vec![0.0f32; tokens * hidden];
2633        for token in 0..tokens {
2634            let row = &qkv[token * 3 * hidden..(token + 1) * 3 * hidden];
2635            value[token * hidden..(token + 1) * hidden]
2636                .copy_from_slice(&row[2 * hidden..3 * hidden]);
2637            for head in 0..heads {
2638                let offset = head * head_dim;
2639                let normed_query = rms_norm(
2640                    &row[offset..offset + head_dim],
2641                    1,
2642                    head_dim,
2643                    query_norm,
2644                    eps,
2645                );
2646                let normed_key = rms_norm(
2647                    &row[hidden + offset..hidden + offset + head_dim],
2648                    1,
2649                    head_dim,
2650                    key_norm,
2651                    eps,
2652                );
2653                let destination = token * hidden + offset;
2654                for dim in 0..half {
2655                    let cos = rope_cos[token * half + dim];
2656                    let sin = rope_sin[token * half + dim];
2657                    let (query_a, query_b) = (normed_query[dim], normed_query[dim + half]);
2658                    query[destination + dim] = query_a * cos - query_b * sin;
2659                    query[destination + dim + half] = query_b * cos + query_a * sin;
2660                    let (key_a, key_b) = (normed_key[dim], normed_key[dim + half]);
2661                    key[destination + dim] = key_a * cos - key_b * sin;
2662                    key[destination + dim + half] = key_b * cos + key_a * sin;
2663                }
2664            }
2665        }
2666        let scale = 1.0 / (head_dim as f32).sqrt();
2667        let mut attended = vec![0.0f32; tokens * hidden];
2668        for token in 0..tokens {
2669            for head in 0..heads {
2670                let mut scores = Vec::with_capacity(tokens);
2671                for source in 0..tokens {
2672                    let mut score = 0.0f32;
2673                    for dim in 0..head_dim {
2674                        score += query[token * hidden + head * head_dim + dim]
2675                            * key[source * hidden + head * head_dim + dim];
2676                    }
2677                    scores.push(score * scale);
2678                }
2679                softmax_in_place(&mut scores);
2680                for (source, probability) in scores.into_iter().enumerate() {
2681                    for dim in 0..head_dim {
2682                        attended[token * hidden + head * head_dim + dim] +=
2683                            probability * value[source * hidden + head * head_dim + dim];
2684                    }
2685                }
2686            }
2687        }
2688        let mut attention = linear(
2689            &attended,
2690            layer_tensor(VisionTensor::AttentionOutput, &[hidden, hidden])?,
2691            tokens,
2692            hidden,
2693            hidden,
2694        );
2695        let attention_bias = layer_tensor(VisionTensor::AttentionOutputBias, &[hidden])?;
2696        for row in attention.chunks_exact_mut(hidden) {
2697            add_in_place(row, attention_bias);
2698        }
2699        add_in_place(&mut x, &attention);
2700        // mlp: rms(norm2) -> gate+bias (max-clamp) / up+bias (+/- clamp) -> silu(g)*u -> down+bias
2701        let mlp_input = rms_norm(
2702            &x,
2703            tokens,
2704            hidden,
2705            layer_tensor(VisionTensor::PreMlpNorm, &[hidden])?,
2706            eps,
2707        );
2708        let mut gate = linear(
2709            &mlp_input,
2710            layer_tensor(VisionTensor::MlpGate, &[ff, hidden])?,
2711            tokens,
2712            hidden,
2713            ff,
2714        );
2715        let gate_bias = layer_tensor(VisionTensor::MlpGateBias, &[ff])?;
2716        let mut up = linear(
2717            &mlp_input,
2718            layer_tensor(VisionTensor::MlpUp, &[ff, hidden])?,
2719            tokens,
2720            hidden,
2721            ff,
2722        );
2723        let up_bias = layer_tensor(VisionTensor::MlpUpBias, &[ff])?;
2724        for row in 0..tokens {
2725            for column in 0..ff {
2726                let index = row * ff + column;
2727                let gated = (gate[index] + gate_bias[column]).min(limit);
2728                let carried = (up[index] + up_bias[column]).clamp(-limit, limit);
2729                gate[index] = silu(gated) * carried;
2730            }
2731        }
2732        let _ = up.drain(..);
2733        let mut down = linear(
2734            &gate,
2735            layer_tensor(VisionTensor::MlpDown, &[hidden, ff])?,
2736            tokens,
2737            ff,
2738            hidden,
2739        );
2740        let down_bias = layer_tensor(VisionTensor::MlpDownBias, &[hidden])?;
2741        for row in down.chunks_exact_mut(hidden) {
2742            add_in_place(row, down_bias);
2743        }
2744        add_in_place(&mut x, &down);
2745    }
2746    let encoder_hidden = rms_norm(
2747        &x,
2748        tokens,
2749        hidden,
2750        tensor(weights, &id(None, VisionTensor::PostEncoderNorm), &[hidden])?,
2751        eps,
2752    );
2753    // Downsample: block-major token groups of merge^2 form the conv2d input
2754    // [hidden, merge, merge]; group rows are (in_row, in_col) row-major by construction.
2755    let downsample_weight =
2756        tensor_5d(VisionTensor::Downsample, &[out_width, hidden, merge, merge])?;
2757    let downsample_bias = tensor(
2758        weights,
2759        &id(None, VisionTensor::DownsampleBias),
2760        &[out_width],
2761    )?;
2762    let groups = tokens / merge_area;
2763    let mut pooled_hidden = vec![0.0f32; groups * out_width];
2764    for group in 0..groups {
2765        for out in 0..out_width {
2766            let mut sum = downsample_bias[out];
2767            for channel in 0..hidden {
2768                for kernel_row in 0..merge {
2769                    for kernel_col in 0..merge {
2770                        let token = group * merge_area + kernel_row * merge + kernel_col;
2771                        sum += downsample_weight
2772                            [((out * hidden + channel) * merge + kernel_row) * merge + kernel_col]
2773                            * encoder_hidden[token * hidden + channel];
2774                    }
2775                }
2776            }
2777            pooled_hidden[group * out_width + out] = sum;
2778        }
2779    }
2780    // Merger: proj (no bias) -> LayerNorm (weight + bias, torch nn.LayerNorm default
2781    // eps 1e-5) -> exact-erf GELU -> clamped gate/up -> silu(gate)*up -> down.
2782    let mut merged = linear(
2783        &pooled_hidden,
2784        tensor(
2785            weights,
2786            &id(None, VisionTensor::MergerProjection),
2787            &[out_width, out_width],
2788        )?,
2789        groups,
2790        out_width,
2791        out_width,
2792    );
2793    let norm_weight = tensor(
2794        weights,
2795        &id(None, VisionTensor::MergerPostProjectionNorm),
2796        &[out_width],
2797    )?;
2798    let norm_bias = tensor(
2799        weights,
2800        &id(None, VisionTensor::MergerPostProjectionNormBias),
2801        &[out_width],
2802    )?;
2803    const LAYER_NORM_EPS: f32 = 1e-5; // torch nn.LayerNorm default (upstream passes none)
2804    for row in merged.chunks_exact_mut(out_width) {
2805        let mean = row.iter().sum::<f32>() / out_width as f32;
2806        let variance = row
2807            .iter()
2808            .map(|value| (value - mean) * (value - mean))
2809            .sum::<f32>()
2810            / out_width as f32;
2811        let inverse = 1.0 / (variance + LAYER_NORM_EPS).sqrt();
2812        for (column, value) in row.iter_mut().enumerate() {
2813            *value = gelu_erf((*value - mean) * inverse * norm_weight[column] + norm_bias[column]);
2814        }
2815    }
2816    let mut merger_gate = linear(
2817        &merged,
2818        tensor(
2819            weights,
2820            &id(None, VisionTensor::MergerGate),
2821            &[proj_inter, out_width],
2822        )?,
2823        groups,
2824        out_width,
2825        proj_inter,
2826    );
2827    let merger_up = linear(
2828        &merged,
2829        tensor(
2830            weights,
2831            &id(None, VisionTensor::MergerUp),
2832            &[proj_inter, out_width],
2833        )?,
2834        groups,
2835        out_width,
2836        proj_inter,
2837    );
2838    for (gate_value, up_value) in merger_gate.iter_mut().zip(merger_up.iter()) {
2839        *gate_value = silu(gate_value.min(limit)) * up_value.clamp(-limit, limit);
2840    }
2841    let projected_hidden = linear(
2842        &merger_gate,
2843        tensor(
2844            weights,
2845            &id(None, VisionTensor::MergerDown),
2846            &[out_width, proj_inter],
2847        )?,
2848        groups,
2849        proj_inter,
2850        out_width,
2851    );
2852    Ok(ReferenceVisionOutput {
2853        encoder_hidden,
2854        pooled_hidden,
2855        projected_hidden,
2856        patch_count: tokens,
2857        output_tokens: groups,
2858        hidden_size: hidden,
2859        projection_size: out_width,
2860    })
2861}
2862
2863#[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
2864fn execute_vision_layer(
2865    plan: &memra_gguf::model_plan::VisionLayerPlan,
2866    weights: &ReferenceWeights,
2867    input: &[f32],
2868    positions: &[[u32; 2]],
2869    tokens: usize,
2870    hidden: usize,
2871) -> Result<Vec<f32>, ReferenceError> {
2872    let id = |tensor| TensorId::Vision {
2873        layer: Some(plan.index),
2874        tensor,
2875    };
2876    let attention_input = rms_norm(
2877        input,
2878        tokens,
2879        hidden,
2880        tensor(weights, &id(VisionTensor::InputNorm), &[hidden])?,
2881        plan.input_norm.epsilon,
2882    );
2883    let query_heads = plan.attention.query_heads as usize;
2884    let kv_heads = plan.attention.kv_heads as usize;
2885    let head_dim = plan.attention.head_dim as usize;
2886    if query_heads == 0 || kv_heads == 0 || query_heads % kv_heads != 0 {
2887        return Err(ReferenceError::InvalidPlan {
2888            layer: Some(plan.index),
2889            reason: "vision attention has invalid query/KV head grouping",
2890        });
2891    }
2892    let mut query = linear(
2893        &attention_input,
2894        tensor(
2895            weights,
2896            &id(VisionTensor::Query),
2897            &[query_heads * head_dim, hidden],
2898        )?,
2899        tokens,
2900        hidden,
2901        query_heads * head_dim,
2902    );
2903    let mut key = linear(
2904        &attention_input,
2905        tensor(
2906            weights,
2907            &id(VisionTensor::Key),
2908            &[kv_heads * head_dim, hidden],
2909        )?,
2910        tokens,
2911        hidden,
2912        kv_heads * head_dim,
2913    );
2914    let mut value = linear(
2915        &attention_input,
2916        tensor(
2917            weights,
2918            &id(VisionTensor::Value),
2919            &[kv_heads * head_dim, hidden],
2920        )?,
2921        tokens,
2922        hidden,
2923        kv_heads * head_dim,
2924    );
2925    apply_optional_head_norm(
2926        weights,
2927        id(VisionTensor::QueryNorm),
2928        &mut query,
2929        tokens * query_heads,
2930        head_dim,
2931        memra_gguf::model_plan::TensorPresence::Required,
2932        plan.input_norm.epsilon,
2933    )?;
2934    apply_optional_head_norm(
2935        weights,
2936        id(VisionTensor::KeyNorm),
2937        &mut key,
2938        tokens * kv_heads,
2939        head_dim,
2940        memra_gguf::model_plan::TensorPresence::Required,
2941        plan.input_norm.epsilon,
2942    )?;
2943    value = rms_norm(
2944        &value,
2945        tokens * kv_heads,
2946        head_dim,
2947        &vec![1.0; head_dim],
2948        plan.input_norm.epsilon,
2949    );
2950    apply_vision_rope(
2951        &mut query,
2952        tokens,
2953        query_heads,
2954        head_dim,
2955        positions,
2956        plan.attention.rope.base,
2957    )?;
2958    apply_vision_rope(
2959        &mut key,
2960        tokens,
2961        kv_heads,
2962        head_dim,
2963        positions,
2964        plan.attention.rope.base,
2965    )?;
2966    let repeat = query_heads / kv_heads;
2967    let mut attended = vec![0.0; tokens * query_heads * head_dim];
2968    for token in 0..tokens {
2969        for head in 0..query_heads {
2970            let kv_head = head / repeat;
2971            let mut scores = Vec::with_capacity(tokens);
2972            for source in 0..tokens {
2973                let mut score = 0.0;
2974                for column in 0..head_dim {
2975                    score += query[(token * query_heads + head) * head_dim + column]
2976                        * key[(source * kv_heads + kv_head) * head_dim + column];
2977                }
2978                scores.push(score);
2979            }
2980            softmax_in_place(&mut scores);
2981            for (source, probability) in scores.into_iter().enumerate() {
2982                for column in 0..head_dim {
2983                    attended[(token * query_heads + head) * head_dim + column] +=
2984                        probability * value[(source * kv_heads + kv_head) * head_dim + column];
2985                }
2986            }
2987        }
2988    }
2989    let attention = linear(
2990        &attended,
2991        tensor(
2992            weights,
2993            &id(VisionTensor::AttentionOutput),
2994            &[hidden, query_heads * head_dim],
2995        )?,
2996        tokens,
2997        query_heads * head_dim,
2998        hidden,
2999    );
3000    let attention = rms_norm(
3001        &attention,
3002        tokens,
3003        hidden,
3004        tensor(weights, &id(VisionTensor::PostAttentionNorm), &[hidden])?,
3005        plan.post_attention_norm.epsilon,
3006    );
3007    let mut residual = input.to_vec();
3008    add_in_place(&mut residual, &attention);
3009    let mlp_input = rms_norm(
3010        &residual,
3011        tokens,
3012        hidden,
3013        tensor(weights, &id(VisionTensor::PreMlpNorm), &[hidden])?,
3014        plan.pre_mlp_norm.epsilon,
3015    );
3016    let intermediate = plan.mlp.intermediate_size as usize;
3017    let gate = linear(
3018        &mlp_input,
3019        tensor(weights, &id(VisionTensor::MlpGate), &[intermediate, hidden])?,
3020        tokens,
3021        hidden,
3022        intermediate,
3023    );
3024    let up = linear(
3025        &mlp_input,
3026        tensor(weights, &id(VisionTensor::MlpUp), &[intermediate, hidden])?,
3027        tokens,
3028        hidden,
3029        intermediate,
3030    );
3031    let mut activated = vec![0.0; gate.len()];
3032    for index in 0..activated.len() {
3033        activated[index] = activate_pair(&plan.mlp.activation, gate[index], up[index], plan.index)?;
3034    }
3035    let mlp = linear(
3036        &activated,
3037        tensor(weights, &id(VisionTensor::MlpDown), &[hidden, intermediate])?,
3038        tokens,
3039        intermediate,
3040        hidden,
3041    );
3042    let mlp = rms_norm(
3043        &mlp,
3044        tokens,
3045        hidden,
3046        tensor(weights, &id(VisionTensor::PostMlpNorm), &[hidden])?,
3047        plan.post_mlp_norm.epsilon,
3048    );
3049    add_in_place(&mut residual, &mlp);
3050    Ok(residual)
3051}
3052
3053#[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
3054fn apply_vision_rope(
3055    values: &mut [f32],
3056    tokens: usize,
3057    heads: usize,
3058    head_dim: usize,
3059    positions: &[[u32; 2]],
3060    base: f32,
3061) -> Result<(), ReferenceError> {
3062    let axes = 2;
3063    let chunk = head_dim / axes;
3064    if head_dim % axes != 0 || !chunk.is_multiple_of(2) || positions.len() != tokens {
3065        return Err(ReferenceError::InvalidPlan {
3066            layer: None,
3067            reason: "vision 2D RoPE requires even per-axis head chunks",
3068        });
3069    }
3070    let half = chunk / 2;
3071    #[allow(clippy::needless_range_loop)]
3072    // allow: the explicit index loop keeps the offset arithmetic visible and aligned with the device-side indexing
3073    for token in 0..tokens {
3074        for head in 0..heads {
3075            let row = (token * heads + head) * head_dim;
3076            #[allow(clippy::needless_range_loop)]
3077            // allow: the explicit index loop keeps the offset arithmetic visible and aligned with the device-side indexing
3078            for axis in 0..axes {
3079                let start = row + axis * chunk;
3080                let position = positions[token][axis] as f32;
3081                for pair in 0..half {
3082                    let angle = position / base.powf((2 * pair) as f32 / chunk as f32);
3083                    let (sin, cos) = angle.sin_cos();
3084                    let left = values[start + pair];
3085                    let right = values[start + half + pair];
3086                    values[start + pair] = left * cos - right * sin;
3087                    values[start + half + pair] = left * sin + right * cos;
3088                }
3089            }
3090        }
3091    }
3092    Ok(())
3093}
3094
3095#[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
3096fn vision_pool(
3097    hidden_states: &[f32],
3098    positions: &[[u32; 2]],
3099    patches: usize,
3100    output_tokens: usize,
3101    hidden: usize,
3102) -> Result<Vec<f32>, ReferenceError> {
3103    if patches % output_tokens != 0 {
3104        return Err(ReferenceError::InvalidPlan {
3105            layer: None,
3106            reason: "vision pooling ratio must divide the patch count",
3107        });
3108    }
3109    let area = patches / output_tokens;
3110    let kernel = (area as f32).sqrt() as usize;
3111    if kernel * kernel != area {
3112        return Err(ReferenceError::InvalidPlan {
3113            layer: None,
3114            reason: "vision pooling ratio must be a square kernel",
3115        });
3116    }
3117    let max_x = positions
3118        .iter()
3119        .map(|position| position[0] as usize)
3120        .max()
3121        .unwrap_or(0)
3122        + 1;
3123    let grid_width = max_x / kernel;
3124    let mut output = vec![0.0; output_tokens * hidden];
3125    for patch in 0..patches {
3126        let target = positions[patch][0] as usize / kernel
3127            + grid_width * (positions[patch][1] as usize / kernel);
3128        if target >= output_tokens {
3129            return Err(ReferenceError::InvalidPlan {
3130                layer: None,
3131                reason: "vision patch positions do not fit the pooled grid",
3132            });
3133        }
3134        for column in 0..hidden {
3135            output[target * hidden + column] +=
3136                hidden_states[patch * hidden + column] / area as f32;
3137        }
3138    }
3139    let scale = (hidden as f32).sqrt();
3140    for value in &mut output {
3141        *value *= scale;
3142    }
3143    Ok(output)
3144}
3145
3146fn collapse_stream_mean(
3147    x: &[f32],
3148    tokens: usize,
3149    hidden: usize,
3150    hyper_streams: Option<usize>,
3151) -> Result<Vec<f32>, ReferenceError> {
3152    let Some(streams) = hyper_streams else {
3153        if x.len() != tokens * hidden {
3154            return Err(ReferenceError::InvalidPlan {
3155                layer: None,
3156                reason: "single-stream DSpark tap has invalid shape",
3157            });
3158        }
3159        return Ok(x.to_vec());
3160    };
3161    if x.len() != tokens * streams * hidden {
3162        return Err(ReferenceError::InvalidPlan {
3163            layer: None,
3164            reason: "HyperConnections DSpark tap has invalid shape",
3165        });
3166    }
3167    let mut output = vec![0.0; tokens * hidden];
3168    for token in 0..tokens {
3169        for stream in 0..streams {
3170            for column in 0..hidden {
3171                output[token * hidden + column] +=
3172                    x[(token * streams + stream) * hidden + column] / streams as f32;
3173            }
3174        }
3175    }
3176    Ok(output)
3177}
3178
3179#[allow(clippy::too_many_arguments)]
3180fn execute_dspark(
3181    plan: &memra_gguf::model_plan::DsparkPlan,
3182    weights: &ReferenceWeights,
3183    token_ids: &[u32],
3184    embedding: &[f32],
3185    output_projection: &[f32],
3186    logits_transforms: &[LogitsTransform],
3187    norm_epsilon: f32,
3188    hidden: usize,
3189    vocab: usize,
3190    taps: Vec<Option<Vec<f32>>>,
3191) -> Result<ReferenceDraftOutput, ReferenceError> {
3192    use memra_gguf::dsv4_forward::{hc_expand, hc_head, matmul, rmsnorm};
3193
3194    let tokens = token_ids.len();
3195    let block_size = plan.block_size as usize;
3196    let rank = plan.markov_rank as usize;
3197    if tokens < 2
3198        || block_size == 0
3199        || plan.blocks.is_empty()
3200        || taps.len() != plan.target_layer_ids.len()
3201        || plan.noise_token_id as usize >= vocab
3202    {
3203        return Err(ReferenceError::InvalidPlan {
3204            layer: None,
3205            reason: "DSpark execution requires a primed prompt and valid drafter geometry",
3206        });
3207    }
3208    let streams = match plan.blocks[0].residual {
3209        ResidualTopology::HyperConnections { streams, .. } if streams > 0 => streams as usize,
3210        _ => {
3211            return Err(ReferenceError::InvalidPlan {
3212                layer: Some(plan.blocks[0].index),
3213                reason: "DSpark blocks require HyperConnections",
3214            });
3215        }
3216    };
3217    let mut main_hidden = vec![0.0; tokens * taps.len() * hidden];
3218    for (target, tap) in taps.into_iter().enumerate() {
3219        let Some(tap) = tap else {
3220            return Err(ReferenceError::InvalidPlan {
3221                layer: None,
3222                reason: "DSpark target layer was not captured from the trunk",
3223            });
3224        };
3225        if tap.len() != tokens * hidden {
3226            return Err(ReferenceError::InvalidPlan {
3227                layer: None,
3228                reason: "DSpark trunk tap has invalid shape",
3229            });
3230        }
3231        for token in 0..tokens {
3232            main_hidden[(token * plan.target_layer_ids.len() + target) * hidden
3233                ..(token * plan.target_layer_ids.len() + target + 1) * hidden]
3234                .copy_from_slice(&tap[token * hidden..(token + 1) * hidden]);
3235        }
3236    }
3237    let main_x = rmsnorm(
3238        &matmul(
3239            &main_hidden,
3240            tokens,
3241            plan.target_layer_ids.len() * hidden,
3242            tensor(
3243                weights,
3244                &TensorId::Dspark(DsparkTensor::MainProjection),
3245                &[hidden, plan.target_layer_ids.len() * hidden],
3246            )?,
3247            hidden,
3248        ),
3249        tensor(
3250            weights,
3251            &TensorId::Dspark(DsparkTensor::MainNorm),
3252            &[hidden],
3253        )?,
3254        norm_epsilon,
3255    );
3256    let rings = plan
3257        .blocks
3258        .iter()
3259        .map(|block| dspark_prime_ring(block, weights, &main_x, tokens, hidden, norm_epsilon))
3260        .collect::<Result<Vec<_>, _>>()?;
3261
3262    let input_token = *token_ids.last().unwrap();
3263    let mut draft_ids = vec![plan.noise_token_id; block_size];
3264    draft_ids[0] = input_token;
3265    let mut embedded = vec![0.0; block_size * hidden];
3266    for (position, &token) in draft_ids.iter().enumerate() {
3267        let token = token as usize;
3268        embedded[position * hidden..(position + 1) * hidden]
3269            .copy_from_slice(&embedding[token * hidden..(token + 1) * hidden]);
3270    }
3271    let mut draft_hidden = hc_expand(&embedded, block_size, streams, hidden);
3272    for (block, ring) in plan.blocks.iter().zip(&rings) {
3273        draft_hidden = execute_dspark_layer(
3274            block,
3275            weights,
3276            &draft_hidden,
3277            ring,
3278            tokens - 1,
3279            block_size,
3280            hidden,
3281            vocab,
3282        )?;
3283    }
3284    let head_set = memra_gguf::dsv4_forward::HcSet {
3285        rows: streams,
3286        fn_w: tensor(
3287            weights,
3288            &TensorId::Dspark(DsparkTensor::HeadHyperFunction),
3289            &[streams, streams * hidden],
3290        )?
3291        .to_vec(),
3292        base: tensor(
3293            weights,
3294            &TensorId::Dspark(DsparkTensor::HeadHyperBase),
3295            &[streams],
3296        )?
3297        .to_vec(),
3298        scale: tensor(
3299            weights,
3300            &TensorId::Dspark(DsparkTensor::HeadHyperScale),
3301            &[1],
3302        )?
3303        .to_vec(),
3304    };
3305    let hc_epsilon = match plan.blocks[0].residual {
3306        ResidualTopology::HyperConnections { epsilon, .. } => epsilon,
3307        _ => unreachable!(),
3308    };
3309    let collapsed = hc_head(
3310        &draft_hidden,
3311        block_size,
3312        streams,
3313        hidden,
3314        &head_set,
3315        norm_epsilon,
3316        hc_epsilon,
3317    );
3318    let normalized = rmsnorm(
3319        &collapsed,
3320        tensor(
3321            weights,
3322            &TensorId::Dspark(DsparkTensor::OutputNorm),
3323            &[hidden],
3324        )?,
3325        norm_epsilon,
3326    );
3327    let mut logits = matmul(&normalized, block_size, hidden, output_projection, vocab);
3328    apply_logits_transforms(&mut logits, vocab, logits_transforms);
3329
3330    let markov_embedding = tensor(
3331        weights,
3332        &TensorId::Dspark(DsparkTensor::MarkovEmbedding),
3333        &[vocab, rank],
3334    )?;
3335    let markov_output = tensor(
3336        weights,
3337        &TensorId::Dspark(DsparkTensor::MarkovOutput),
3338        &[vocab, rank],
3339    )?;
3340    let confidence_weight = tensor(
3341        weights,
3342        &TensorId::Dspark(DsparkTensor::ConfidenceProjection),
3343        &[1, hidden + rank],
3344    )?;
3345    let mut output_ids = vec![input_token];
3346    let mut confidence = Vec::with_capacity(block_size);
3347    for position in 0..block_size {
3348        let previous = output_ids[position] as usize;
3349        let markov = &markov_embedding[previous * rank..(previous + 1) * rank];
3350        let row = &mut logits[position * vocab..(position + 1) * vocab];
3351        for token in 0..vocab {
3352            row[token] += memra_gguf::dsv4_forward::dot(
3353                markov,
3354                &markov_output[token * rank..(token + 1) * rank],
3355            );
3356        }
3357        let next = row
3358            .iter()
3359            .enumerate()
3360            .max_by(|(left_index, left), (right_index, right)| {
3361                left.total_cmp(right)
3362                    .then_with(|| right_index.cmp(left_index))
3363            })
3364            .map(|(index, _)| index as u32)
3365            .unwrap();
3366        output_ids.push(next);
3367        let mut confidence_input = Vec::with_capacity(hidden + rank);
3368        confidence_input.extend_from_slice(&collapsed[position * hidden..(position + 1) * hidden]);
3369        confidence_input.extend_from_slice(markov);
3370        confidence.push(memra_gguf::dsv4_forward::dot(
3371            &confidence_input,
3372            confidence_weight,
3373        ));
3374    }
3375    Ok(ReferenceDraftOutput {
3376        input_token,
3377        output_ids,
3378        confidence,
3379        logits,
3380        hidden: collapsed,
3381        block_size,
3382    })
3383}
3384
3385fn dspark_prime_ring(
3386    layer: &memra_gguf::model_plan::LayerPlan,
3387    weights: &ReferenceWeights,
3388    main_x: &[f32],
3389    tokens: usize,
3390    hidden: usize,
3391    epsilon: f32,
3392) -> Result<Vec<f32>, ReferenceError> {
3393    use memra_gguf::dsv4_forward::{ActQuantVariant, apply_rope, matmul, rmsnorm};
3394    use memra_gguf::model_plan::{MlaAttentionPlan, RopeFactors, SparseIndexPlan};
3395
3396    let AttentionPlan::Mla(MlaAttentionPlan::CompressedKv {
3397        latent_head_dim,
3398        rope_head_dim,
3399        window,
3400        rope,
3401        compressor: None,
3402        sparse_index: SparseIndexPlan::None,
3403        ..
3404    }) = &layer.attention
3405    else {
3406        return Err(ReferenceError::InvalidPlan {
3407            layer: Some(layer.index),
3408            reason: "DSpark blocks require uncompressed window-only attention",
3409        });
3410    };
3411    if !matches!(rope.factors, RopeFactors::None) {
3412        return Err(ReferenceError::InvalidPlan {
3413            layer: Some(layer.index),
3414            reason: "DSpark block RoPE must not use scaling factors",
3415        });
3416    }
3417    let head_dim = *latent_head_dim as usize;
3418    let rope_dim = *rope_head_dim as usize;
3419    if head_dim <= rope_dim || !(head_dim - rope_dim).is_multiple_of(64) {
3420        return Err(ReferenceError::InvalidPlan {
3421            layer: Some(layer.index),
3422            reason: "DSpark block has invalid KV quantization geometry",
3423        });
3424    }
3425    let frequencies = memra_gguf::dsv4_forward::precompute_freqs_cis(
3426        rope_dim,
3427        tokens + 1,
3428        0,
3429        rope.base,
3430        1.0,
3431        32.0,
3432        1.0,
3433    );
3434    let mut key_value = rmsnorm(
3435        &matmul(
3436            main_x,
3437            tokens,
3438            hidden,
3439            tensor(
3440                weights,
3441                &layer_id(layer.index, LayerTensor::MlaKvDown),
3442                &[head_dim, hidden],
3443            )?,
3444            head_dim,
3445        ),
3446        tensor(
3447            weights,
3448            &layer_id(layer.index, LayerTensor::MlaKvDownNorm),
3449            &[head_dim],
3450        )?,
3451        epsilon,
3452    );
3453    let positions: Vec<_> = (0..tokens).collect();
3454    apply_rope(
3455        &mut key_value,
3456        tokens,
3457        1,
3458        head_dim,
3459        rope_dim,
3460        &frequencies,
3461        &positions,
3462        false,
3463    );
3464    for row in key_value.chunks_exact_mut(head_dim) {
3465        memra_gguf::dsv4_forward::act_quant(
3466            &mut row[..head_dim - rope_dim],
3467            64,
3468            ActQuantVariant::RefFp8Round,
3469        );
3470    }
3471    let window = *window as usize;
3472    let mut ring = vec![0.0; window * head_dim];
3473    for position in tokens.saturating_sub(window)..tokens {
3474        ring[(position % window) * head_dim..(position % window + 1) * head_dim]
3475            .copy_from_slice(&key_value[position * head_dim..(position + 1) * head_dim]);
3476    }
3477    Ok(ring)
3478}
3479
3480#[allow(clippy::too_many_arguments)]
3481fn execute_dspark_layer(
3482    layer: &memra_gguf::model_plan::LayerPlan,
3483    weights: &ReferenceWeights,
3484    input: &[f32],
3485    ring: &[f32],
3486    start_position: usize,
3487    block_size: usize,
3488    hidden: usize,
3489    vocab: usize,
3490) -> Result<Vec<f32>, ReferenceError> {
3491    let ResidualTopology::HyperConnections {
3492        streams,
3493        epsilon,
3494        sinkhorn_iterations,
3495        collapse: _,
3496    } = layer.residual
3497    else {
3498        return Err(ReferenceError::InvalidPlan {
3499            layer: Some(layer.index),
3500            reason: "DSpark block requires HyperConnections",
3501        });
3502    };
3503    let streams = streams as usize;
3504    let attention_set = hyper_set(
3505        weights,
3506        layer.index,
3507        streams,
3508        hidden,
3509        LayerTensor::HyperAttentionFunction,
3510        LayerTensor::HyperAttentionBase,
3511        LayerTensor::HyperAttentionScale,
3512    )?;
3513    let (attention_input, post, combination) = memra_gguf::dsv4_forward::hc_pre(
3514        input,
3515        block_size,
3516        streams,
3517        hidden,
3518        &attention_set,
3519        sinkhorn_iterations,
3520        epsilon,
3521    );
3522    let attention_input = rms_norm(
3523        &attention_input,
3524        block_size,
3525        hidden,
3526        tensor(
3527            weights,
3528            &layer_id(layer.index, LayerTensor::PreAttentionNorm),
3529            &[hidden],
3530        )?,
3531        layer.pre_attention_norm.epsilon,
3532    );
3533    let attention = dspark_attention(
3534        layer,
3535        weights,
3536        &attention_input,
3537        ring,
3538        start_position,
3539        block_size,
3540        hidden,
3541    )?;
3542    let attention_residual = memra_gguf::dsv4_forward::hc_post(
3543        &attention,
3544        input,
3545        block_size,
3546        streams,
3547        hidden,
3548        &post,
3549        &combination,
3550    );
3551    let mlp_set = hyper_set(
3552        weights,
3553        layer.index,
3554        streams,
3555        hidden,
3556        LayerTensor::HyperMlpFunction,
3557        LayerTensor::HyperMlpBase,
3558        LayerTensor::HyperMlpScale,
3559    )?;
3560    let (mlp_input, post, combination) = memra_gguf::dsv4_forward::hc_pre(
3561        &attention_residual,
3562        block_size,
3563        streams,
3564        hidden,
3565        &mlp_set,
3566        sinkhorn_iterations,
3567        epsilon,
3568    );
3569    let mlp_input = rms_norm(
3570        &mlp_input,
3571        block_size,
3572        hidden,
3573        tensor(
3574            weights,
3575            &layer_id(layer.index, LayerTensor::PreMlpNorm),
3576            &[hidden],
3577        )?,
3578        layer.pre_mlp_norm.epsilon,
3579    );
3580    let zeros = vec![0; block_size];
3581    let mlp = match &layer.mlp {
3582        MlpPlan::Dense(mlp) => {
3583            dense_mlp(layer.index, mlp, weights, &mlp_input, block_size, hidden)?
3584        }
3585        MlpPlan::Moe(moe) => moe_mlp(
3586            layer.index,
3587            moe,
3588            weights,
3589            &mlp_input,
3590            &zeros,
3591            block_size,
3592            hidden,
3593            vocab,
3594        )?,
3595    };
3596    Ok(memra_gguf::dsv4_forward::hc_post(
3597        &mlp,
3598        &attention_residual,
3599        block_size,
3600        streams,
3601        hidden,
3602        &post,
3603        &combination,
3604    ))
3605}
3606
3607#[allow(clippy::too_many_arguments)]
3608#[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
3609fn dspark_attention(
3610    layer: &memra_gguf::model_plan::LayerPlan,
3611    weights: &ReferenceWeights,
3612    x: &[f32],
3613    ring: &[f32],
3614    start_position: usize,
3615    block_size: usize,
3616    hidden: usize,
3617) -> Result<Vec<f32>, ReferenceError> {
3618    use memra_gguf::dsv4_forward::{ActQuantVariant, apply_rope, matmul, rmsnorm};
3619    use memra_gguf::model_plan::{MlaAttentionPlan, RopeFactors, SparseIndexPlan};
3620
3621    let AttentionPlan::Mla(MlaAttentionPlan::CompressedKv {
3622        query_heads,
3623        q_lora_rank,
3624        latent_head_dim,
3625        rope_head_dim,
3626        output_lora_rank,
3627        output_groups,
3628        window,
3629        rope,
3630        compressor: None,
3631        sparse_index: SparseIndexPlan::None,
3632    }) = &layer.attention
3633    else {
3634        return Err(ReferenceError::InvalidPlan {
3635            layer: Some(layer.index),
3636            reason: "DSpark block requires window-only compressed-attention geometry",
3637        });
3638    };
3639    if !matches!(rope.factors, RopeFactors::None) {
3640        return Err(ReferenceError::InvalidPlan {
3641            layer: Some(layer.index),
3642            reason: "DSpark block RoPE must not use scaling factors",
3643        });
3644    }
3645    let heads = *query_heads as usize;
3646    let q_rank = *q_lora_rank as usize;
3647    let head_dim = *latent_head_dim as usize;
3648    let rope_dim = *rope_head_dim as usize;
3649    let output_rank = *output_lora_rank as usize;
3650    let groups = *output_groups as usize;
3651    let window = *window as usize;
3652    if start_position == 0
3653        || head_dim <= rope_dim
3654        || !(head_dim - rope_dim).is_multiple_of(64)
3655        || groups == 0
3656        || heads % groups != 0
3657        || ring.len() != window * head_dim
3658    {
3659        return Err(ReferenceError::InvalidPlan {
3660            layer: Some(layer.index),
3661            reason: "DSpark attention has invalid geometry or unprimed ring",
3662        });
3663    }
3664    let positions: Vec<_> = (1..=block_size)
3665        .map(|offset| start_position + offset)
3666        .collect();
3667    let frequencies = memra_gguf::dsv4_forward::precompute_freqs_cis(
3668        rope_dim,
3669        start_position + block_size + 1,
3670        0,
3671        rope.base,
3672        1.0,
3673        32.0,
3674        1.0,
3675    );
3676    let query_low_rank = rmsnorm(
3677        &matmul(
3678            x,
3679            block_size,
3680            hidden,
3681            tensor(
3682                weights,
3683                &layer_id(layer.index, LayerTensor::MlaQueryDown),
3684                &[q_rank, hidden],
3685            )?,
3686            q_rank,
3687        ),
3688        tensor(
3689            weights,
3690            &layer_id(layer.index, LayerTensor::MlaQueryDownNorm),
3691            &[q_rank],
3692        )?,
3693        layer.pre_attention_norm.epsilon,
3694    );
3695    let mut query = matmul(
3696        &query_low_rank,
3697        block_size,
3698        q_rank,
3699        tensor(
3700            weights,
3701            &layer_id(layer.index, LayerTensor::MlaQueryUp),
3702            &[heads * head_dim, q_rank],
3703        )?,
3704        heads * head_dim,
3705    );
3706    for head in query.chunks_exact_mut(head_dim) {
3707        let mean_square = head
3708            .iter()
3709            .map(|value| (*value as f64) * (*value as f64))
3710            .sum::<f64>()
3711            / head_dim as f64;
3712        let scale = 1.0 / (mean_square as f32 + layer.pre_attention_norm.epsilon).sqrt();
3713        for value in head {
3714            *value *= scale;
3715        }
3716    }
3717    apply_rope(
3718        &mut query,
3719        block_size,
3720        heads,
3721        head_dim,
3722        rope_dim,
3723        &frequencies,
3724        &positions,
3725        false,
3726    );
3727    let mut key_value = rmsnorm(
3728        &matmul(
3729            x,
3730            block_size,
3731            hidden,
3732            tensor(
3733                weights,
3734                &layer_id(layer.index, LayerTensor::MlaKvDown),
3735                &[head_dim, hidden],
3736            )?,
3737            head_dim,
3738        ),
3739        tensor(
3740            weights,
3741            &layer_id(layer.index, LayerTensor::MlaKvDownNorm),
3742            &[head_dim],
3743        )?,
3744        layer.pre_attention_norm.epsilon,
3745    );
3746    apply_rope(
3747        &mut key_value,
3748        block_size,
3749        1,
3750        head_dim,
3751        rope_dim,
3752        &frequencies,
3753        &positions,
3754        false,
3755    );
3756    for row in key_value.chunks_exact_mut(head_dim) {
3757        memra_gguf::dsv4_forward::act_quant(
3758            &mut row[..head_dim - rope_dim],
3759            64,
3760            ActQuantVariant::RefFp8Round,
3761        );
3762    }
3763    let indices = memra_gguf::dsv4_dspark::dspark_topk_idxs(window, block_size, start_position);
3764    let sink = tensor(
3765        weights,
3766        &layer_id(layer.index, LayerTensor::AttentionSink),
3767        &[heads],
3768    )?;
3769    let mut attended = vec![0.0; block_size * heads * head_dim];
3770    for token in 0..block_size {
3771        memra_gguf::dsv4_decode::sparse_attn_query(
3772            &query[token * heads * head_dim..(token + 1) * heads * head_dim],
3773            heads,
3774            head_dim,
3775            &indices,
3776            |index| {
3777                if index < window {
3778                    &ring[index * head_dim..(index + 1) * head_dim]
3779                } else {
3780                    let index = index - window;
3781                    &key_value[index * head_dim..(index + 1) * head_dim]
3782                }
3783            },
3784            sink,
3785            (head_dim as f64).powf(-0.5) as f32,
3786            &mut attended[token * heads * head_dim..(token + 1) * heads * head_dim],
3787        );
3788    }
3789    apply_rope(
3790        &mut attended,
3791        block_size,
3792        heads,
3793        head_dim,
3794        rope_dim,
3795        &frequencies,
3796        &positions,
3797        true,
3798    );
3799    let group_width = heads / groups * head_dim;
3800    let output_down = tensor(
3801        weights,
3802        &layer_id(layer.index, LayerTensor::MlaOutputDown),
3803        &[groups * output_rank, group_width],
3804    )?;
3805    let mut grouped = vec![0.0; block_size * groups * output_rank];
3806    for token in 0..block_size {
3807        for group in 0..groups {
3808            let source = &attended[token * heads * head_dim + group * group_width
3809                ..token * heads * head_dim + (group + 1) * group_width];
3810            for rank in 0..output_rank {
3811                let weight = &output_down[(group * output_rank + rank) * group_width
3812                    ..(group * output_rank + rank + 1) * group_width];
3813                grouped[(token * groups + group) * output_rank + rank] =
3814                    memra_gguf::dsv4_forward::dot(source, weight);
3815            }
3816        }
3817    }
3818    Ok(matmul(
3819        &grouped,
3820        block_size,
3821        groups * output_rank,
3822        tensor(
3823            weights,
3824            &layer_id(layer.index, LayerTensor::MlaOutput),
3825            &[hidden, groups * output_rank],
3826        )?,
3827        hidden,
3828    ))
3829}
3830
3831fn hyper_topology(
3832    plan: &ModelPlan,
3833) -> Result<Option<(usize, f32, u32, HcCollapse)>, ReferenceError> {
3834    let topology = plan.layers.iter().find_map(|layer| match layer.residual {
3835        ResidualTopology::HyperConnections {
3836            streams,
3837            epsilon,
3838            sinkhorn_iterations,
3839            collapse,
3840        } => Some((streams as usize, epsilon, sinkhorn_iterations, collapse)),
3841        _ => None,
3842    });
3843    let Some(topology) = topology else {
3844        return Ok(None);
3845    };
3846    if topology.0 == 0 || topology.1 <= 0.0 || topology.2 == 0 {
3847        return Err(ReferenceError::InvalidPlan {
3848            layer: None,
3849            reason: "HyperConnections require streams, epsilon, and Sinkhorn iterations",
3850        });
3851    }
3852    for layer in &plan.layers {
3853        if layer.residual
3854            != (ResidualTopology::HyperConnections {
3855                streams: topology.0 as u32,
3856                epsilon: topology.1,
3857                sinkhorn_iterations: topology.2,
3858                collapse: topology.3,
3859            })
3860        {
3861            return Err(ReferenceError::InvalidPlan {
3862                layer: Some(layer.index),
3863                reason: "HyperConnections topology must be consistent across the trunk",
3864            });
3865        }
3866    }
3867    Ok(Some(topology))
3868}
3869
3870/// qwen4_exp gated-residual topology: `(streams, bottleneck_rank)` when the trunk runs the
3871/// 4-branch wide stream. Requires topology consistency across trunk AND MTP blocks, plus a
3872/// matching exit mixer — the model has no final norm to fall back to (SEMANTICS.md).
3873fn gated_residual_topology(plan: &ModelPlan) -> Result<Option<(usize, usize)>, ReferenceError> {
3874    let topology = plan.layers.iter().find_map(|layer| match layer.residual {
3875        ResidualTopology::GatedResidual {
3876            streams,
3877            bottleneck_rank,
3878        } => Some((streams as usize, bottleneck_rank as usize)),
3879        _ => None,
3880    });
3881    let Some((streams, rank)) = topology else {
3882        if plan.exit_mixer.is_some() {
3883            return Err(ReferenceError::InvalidPlan {
3884                layer: None,
3885                reason: "exit mixer requires a gated-residual trunk",
3886            });
3887        }
3888        return Ok(None);
3889    };
3890    if streams == 0 || rank == 0 {
3891        return Err(ReferenceError::InvalidPlan {
3892            layer: None,
3893            reason: "gated residual requires streams and a bottleneck rank",
3894        });
3895    }
3896    for layer in plan
3897        .layers
3898        .iter()
3899        .chain(plan.mtp_blocks.iter().map(|block| &block.layer))
3900    {
3901        if layer.residual
3902            != (ResidualTopology::GatedResidual {
3903                streams: streams as u32,
3904                bottleneck_rank: rank as u32,
3905            })
3906        {
3907            return Err(ReferenceError::InvalidPlan {
3908                layer: Some(layer.index),
3909                reason: "gated-residual topology must be consistent across trunk and MTP blocks",
3910            });
3911        }
3912    }
3913    match plan.exit_mixer {
3914        Some(mixer)
3915            if mixer.streams as usize == streams && mixer.bottleneck_rank as usize == rank => {}
3916        _ => {
3917            return Err(ReferenceError::InvalidPlan {
3918                layer: None,
3919                reason: "gated-residual trunk requires a matching exit mixer",
3920            });
3921        }
3922    }
3923    Ok(Some((streams, rank)))
3924}
3925
3926fn collapse_hyper_head(
3927    weights: &ReferenceWeights,
3928    x: &[f32],
3929    tokens: usize,
3930    streams: usize,
3931    hidden: usize,
3932    plan: &ModelPlan,
3933    epsilon: f32,
3934) -> Result<Vec<f32>, ReferenceError> {
3935    let set = memra_gguf::dsv4_forward::HcSet {
3936        rows: streams,
3937        fn_w: tensor(
3938            weights,
3939            &TensorId::HyperHeadFunction,
3940            &[streams, streams * hidden],
3941        )?
3942        .to_vec(),
3943        base: tensor(weights, &TensorId::HyperHeadBase, &[streams])?.to_vec(),
3944        scale: tensor(weights, &TensorId::HyperHeadScale, &[1])?.to_vec(),
3945    };
3946    Ok(memra_gguf::dsv4_forward::hc_head(
3947        x,
3948        tokens,
3949        streams,
3950        hidden,
3951        &set,
3952        plan.output_norm.epsilon,
3953        epsilon,
3954    ))
3955}
3956
3957fn apply_logits_transforms(logits: &mut [f32], vocab: usize, transforms: &[LogitsTransform]) {
3958    for transform in transforms {
3959        match transform {
3960            LogitsTransform::Softcap(cap) => {
3961                for value in logits.iter_mut() {
3962                    *value = *cap * (*value / *cap).tanh();
3963                }
3964            }
3965            LogitsTransform::SuppressTokens(ids) => {
3966                for row in logits.chunks_exact_mut(vocab) {
3967                    for &id in ids {
3968                        if let Some(value) = row.get_mut(id as usize) {
3969                            *value = f32::NEG_INFINITY;
3970                        }
3971                    }
3972                }
3973            }
3974        }
3975    }
3976}
3977
3978#[allow(clippy::too_many_arguments)]
3979fn execute_layer(
3980    layer: &memra_gguf::model_plan::LayerPlan,
3981    weights: &ReferenceWeights,
3982    input: &[f32],
3983    token_ids: &[u32],
3984    tokens: usize,
3985    hidden: usize,
3986    vocab: usize,
3987    scope: LayerScope,
3988) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
3989    if let ResidualTopology::GatedResidual {
3990        streams,
3991        bottleneck_rank,
3992    } = layer.residual
3993    {
3994        return execute_gated_residual_layer(
3995            layer,
3996            weights,
3997            input,
3998            token_ids,
3999            tokens,
4000            hidden,
4001            vocab,
4002            streams as usize,
4003            bottleneck_rank as usize,
4004            scope,
4005        );
4006    }
4007    // The QSA overlay and PLE are programs of the gated-residual layer; running them
4008    // through any other residual arm would silently drop them.
4009    if layer.sparse_overlay.is_some() || layer.ple.is_some() {
4010        return Err(ReferenceError::UnsupportedOperation {
4011            layer: Some(layer.index),
4012            operation: "sparse overlay / PLE outside the gated-residual program",
4013        });
4014    }
4015    if let ResidualTopology::HyperConnections {
4016        streams,
4017        epsilon,
4018        sinkhorn_iterations,
4019        // The collapse knob only applies at model exit; per-layer mixing is identical.
4020        collapse: _,
4021    } = layer.residual
4022    {
4023        return execute_hyper_layer(
4024            layer,
4025            weights,
4026            input,
4027            token_ids,
4028            tokens,
4029            hidden,
4030            vocab,
4031            streams as usize,
4032            epsilon,
4033            sinkhorn_iterations,
4034        );
4035    }
4036    if let ResidualTopology::Gemma {
4037        parallel_moe: Some(parallel),
4038        ..
4039    } = layer.residual
4040    {
4041        return execute_gemma_parallel_moe_layer(layer, parallel, weights, input, tokens, hidden);
4042    }
4043    if let ResidualTopology::Gemma {
4044        parallel_moe: None, ..
4045    } = layer.residual
4046    {
4047        return execute_gemma_dense_layer(layer, weights, input, tokens, hidden);
4048    }
4049    if layer.residual != ResidualTopology::Serial {
4050        return Err(ReferenceError::UnsupportedOperation {
4051            layer: Some(layer.index),
4052            operation: "non-serial residual",
4053        });
4054    }
4055    let pre_attn = rms_norm(
4056        input,
4057        tokens,
4058        hidden,
4059        tensor(
4060            weights,
4061            &layer_id(layer.index, LayerTensor::PreAttentionNorm),
4062            &[hidden],
4063        )?,
4064        layer.pre_attention_norm.epsilon,
4065    );
4066    let (attention, layer_state) = match &layer.attention {
4067        AttentionPlan::Full(attention) => full_attention(
4068            layer.index,
4069            attention,
4070            None,
4071            layer.pre_attention_norm.epsilon,
4072            weights,
4073            &pre_attn,
4074            tokens,
4075            hidden,
4076            None,
4077        )?,
4078        AttentionPlan::SlidingWindow { attention, window } => full_attention(
4079            layer.index,
4080            attention,
4081            Some(*window as usize),
4082            layer.pre_attention_norm.epsilon,
4083            weights,
4084            &pre_attn,
4085            tokens,
4086            hidden,
4087            None,
4088        )?,
4089        AttentionPlan::Mla(mla) => mla_attention(
4090            layer.index,
4091            mla,
4092            layer.pre_attention_norm.epsilon,
4093            weights,
4094            &pre_attn,
4095            tokens,
4096            hidden,
4097        )?,
4098        AttentionPlan::GatedDeltaNet(gdn) => gated_delta_net(
4099            layer.index,
4100            gdn,
4101            layer.pre_attention_norm.epsilon,
4102            weights,
4103            &pre_attn,
4104            tokens,
4105            hidden,
4106        )?,
4107        AttentionPlan::KimiDeltaNet(kda) => kimi_delta_net(
4108            layer.index,
4109            kda,
4110            layer.pre_attention_norm.epsilon,
4111            weights,
4112            &pre_attn,
4113            tokens,
4114            hidden,
4115        )?,
4116    };
4117    let mut output = input.to_vec();
4118    add_in_place(&mut output, &attention);
4119    let pre_mlp = rms_norm(
4120        &output,
4121        tokens,
4122        hidden,
4123        tensor(
4124            weights,
4125            &layer_id(layer.index, LayerTensor::PreMlpNorm),
4126            &[hidden],
4127        )?,
4128        layer.pre_mlp_norm.epsilon,
4129    );
4130    let mlp = match &layer.mlp {
4131        MlpPlan::Dense(mlp) => dense_mlp(layer.index, mlp, weights, &pre_mlp, tokens, hidden)?,
4132        MlpPlan::Moe(moe) => moe_mlp(
4133            layer.index,
4134            moe,
4135            weights,
4136            &pre_mlp,
4137            token_ids,
4138            tokens,
4139            hidden,
4140            vocab,
4141        )?,
4142    };
4143    add_in_place(&mut output, &mlp);
4144    Ok((output, layer_state))
4145}
4146
4147fn execute_gemma_parallel_moe_layer(
4148    layer: &memra_gguf::model_plan::LayerPlan,
4149    parallel: memra_gguf::model_plan::GemmaParallelMoePlan,
4150    weights: &ReferenceWeights,
4151    input: &[f32],
4152    tokens: usize,
4153    hidden: usize,
4154) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
4155    let ResidualTopology::Gemma {
4156        post_attention_norm,
4157        post_mlp_norm,
4158        layer_scale,
4159        parallel_moe: Some(_),
4160    } = layer.residual
4161    else {
4162        unreachable!()
4163    };
4164    let pre_attention = rms_norm(
4165        input,
4166        tokens,
4167        hidden,
4168        tensor(
4169            weights,
4170            &layer_id(layer.index, LayerTensor::PreAttentionNorm),
4171            &[hidden],
4172        )?,
4173        layer.pre_attention_norm.epsilon,
4174    );
4175    let (attention, state) = match &layer.attention {
4176        AttentionPlan::Full(attention) => full_attention(
4177            layer.index,
4178            attention,
4179            None,
4180            layer.pre_attention_norm.epsilon,
4181            weights,
4182            &pre_attention,
4183            tokens,
4184            hidden,
4185            None,
4186        )?,
4187        AttentionPlan::SlidingWindow { attention, window } => full_attention(
4188            layer.index,
4189            attention,
4190            Some(*window as usize),
4191            layer.pre_attention_norm.epsilon,
4192            weights,
4193            &pre_attention,
4194            tokens,
4195            hidden,
4196            None,
4197        )?,
4198        _ => {
4199            return Err(ReferenceError::UnsupportedOperation {
4200                layer: Some(layer.index),
4201                operation: "gemma parallel MoE non-softmax attention",
4202            });
4203        }
4204    };
4205    let attention = rms_norm(
4206        &attention,
4207        tokens,
4208        hidden,
4209        tensor(
4210            weights,
4211            &layer_id(layer.index, LayerTensor::PostAttentionNorm),
4212            &[hidden],
4213        )?,
4214        post_attention_norm.epsilon,
4215    );
4216    let mut attention_residual = input.to_vec();
4217    add_in_place(&mut attention_residual, &attention);
4218
4219    let MlpPlan::Moe(moe) = &layer.mlp else {
4220        return Err(ReferenceError::InvalidPlan {
4221            layer: Some(layer.index),
4222            reason: "gemma parallel MoE residual requires an MoE plan",
4223        });
4224    };
4225    let shared_plan = moe.shared.as_ref().ok_or(ReferenceError::InvalidPlan {
4226        layer: Some(layer.index),
4227        reason: "gemma parallel MoE requires a shared MLP branch",
4228    })?;
4229    let shared_input = rms_norm(
4230        &attention_residual,
4231        tokens,
4232        hidden,
4233        tensor(
4234            weights,
4235            &layer_id(layer.index, LayerTensor::PreMlpNorm),
4236            &[hidden],
4237        )?,
4238        layer.pre_mlp_norm.epsilon,
4239    );
4240    let shared_intermediate = shared_plan.intermediate_size as usize;
4241    let shared_gate = linear(
4242        &shared_input,
4243        tensor(
4244            weights,
4245            &layer_id(layer.index, LayerTensor::SharedMlpGate),
4246            &[shared_intermediate, hidden],
4247        )?,
4248        tokens,
4249        hidden,
4250        shared_intermediate,
4251    );
4252    let shared_up = linear(
4253        &shared_input,
4254        tensor(
4255            weights,
4256            &layer_id(layer.index, LayerTensor::SharedMlpUp),
4257            &[shared_intermediate, hidden],
4258        )?,
4259        tokens,
4260        hidden,
4261        shared_intermediate,
4262    );
4263    let mut shared_activated = vec![0.0; shared_gate.len()];
4264    for index in 0..shared_activated.len() {
4265        shared_activated[index] = activate_pair(
4266            &moe.activation,
4267            shared_gate[index],
4268            shared_up[index],
4269            layer.index,
4270        )?;
4271    }
4272    let shared = linear(
4273        &shared_activated,
4274        tensor(
4275            weights,
4276            &layer_id(layer.index, LayerTensor::SharedMlpDown),
4277            &[hidden, shared_intermediate],
4278        )?,
4279        tokens,
4280        shared_intermediate,
4281        hidden,
4282    );
4283    let shared = rms_norm(
4284        &shared,
4285        tokens,
4286        hidden,
4287        tensor(
4288            weights,
4289            &layer_id(layer.index, LayerTensor::PostSharedMlpNorm),
4290            &[hidden],
4291        )?,
4292        parallel.shared_post_norm.epsilon,
4293    );
4294
4295    let routed_input = rms_norm(
4296        &attention_residual,
4297        tokens,
4298        hidden,
4299        tensor(
4300            weights,
4301            &layer_id(layer.index, LayerTensor::PreRoutedMlpNorm),
4302            &[hidden],
4303        )?,
4304        parallel.routed_pre_norm.epsilon,
4305    );
4306    let router_scale = tensor(
4307        weights,
4308        &layer_id(layer.index, LayerTensor::MoeRouterScale),
4309        &[hidden],
4310    )?;
4311    let router_weight: Vec<_> = router_scale
4312        .iter()
4313        .map(|value| *value / (hidden as f32).sqrt())
4314        .collect();
4315    let router_input = rms_norm(
4316        &attention_residual,
4317        tokens,
4318        hidden,
4319        &router_weight,
4320        layer.pre_mlp_norm.epsilon,
4321    );
4322    let experts = moe.expert_count as usize;
4323    let selected = moe.experts_per_token as usize;
4324    let intermediate = moe.expert_intermediate_size as usize;
4325    let router_logits = linear(
4326        &router_input,
4327        tensor(
4328            weights,
4329            &layer_id(layer.index, LayerTensor::MoeRouter),
4330            &[experts, hidden],
4331        )?,
4332        tokens,
4333        hidden,
4334        experts,
4335    );
4336    let gate_up = tensor(
4337        weights,
4338        &layer_id(layer.index, LayerTensor::MoeExpertGateUpBank),
4339        &[experts, 2 * intermediate, hidden],
4340    )?;
4341    let down = tensor(
4342        weights,
4343        &layer_id(layer.index, LayerTensor::MoeExpertDownBank),
4344        &[experts, hidden, intermediate],
4345    )?;
4346    let expert_scale = tensor(
4347        weights,
4348        &layer_id(layer.index, LayerTensor::MoeExpertOutputScale),
4349        &[experts],
4350    )?;
4351    let mut routed = vec![0.0; tokens * hidden];
4352    for token in 0..tokens {
4353        let routes = route_experts(
4354            &moe.router,
4355            &router_logits[token * experts..(token + 1) * experts],
4356            None,
4357            selected,
4358            None,
4359            layer.index,
4360        )?;
4361        let row = &routed_input[token * hidden..(token + 1) * hidden];
4362        for (expert, route_weight) in routes {
4363            let expert_offset = expert * 2 * intermediate * hidden;
4364            let mut activated = vec![0.0; intermediate];
4365            for output in 0..intermediate {
4366                let gate = memra_gguf::dsv4_forward::dot(
4367                    row,
4368                    &gate_up
4369                        [expert_offset + output * hidden..expert_offset + (output + 1) * hidden],
4370                );
4371                let up_offset = expert_offset + (intermediate + output) * hidden;
4372                let up =
4373                    memra_gguf::dsv4_forward::dot(row, &gate_up[up_offset..up_offset + hidden]);
4374                activated[output] = activate_pair(&moe.activation, gate, up, layer.index)?;
4375            }
4376            let down_offset = expert * hidden * intermediate;
4377            for output in 0..hidden {
4378                routed[token * hidden + output] += route_weight
4379                    * expert_scale[expert]
4380                    * memra_gguf::dsv4_forward::dot(
4381                        &activated,
4382                        &down[down_offset + output * intermediate
4383                            ..down_offset + (output + 1) * intermediate],
4384                    );
4385            }
4386        }
4387    }
4388    let routed = rms_norm(
4389        &routed,
4390        tokens,
4391        hidden,
4392        tensor(
4393            weights,
4394            &layer_id(layer.index, LayerTensor::PostRoutedMlpNorm),
4395            &[hidden],
4396        )?,
4397        parallel.routed_post_norm.epsilon,
4398    );
4399    let mut combined = shared;
4400    add_in_place(&mut combined, &routed);
4401    let combined = rms_norm(
4402        &combined,
4403        tokens,
4404        hidden,
4405        tensor(
4406            weights,
4407            &layer_id(layer.index, LayerTensor::PostMlpNorm),
4408            &[hidden],
4409        )?,
4410        post_mlp_norm.epsilon,
4411    );
4412    add_in_place(&mut attention_residual, &combined);
4413    let scale = match layer_scale {
4414        GemmaLayerScale::Learned => tensor(
4415            weights,
4416            &layer_id(layer.index, LayerTensor::LayerScale),
4417            &[1],
4418        )?[0],
4419    };
4420    for value in &mut attention_residual {
4421        *value *= scale;
4422    }
4423    Ok((attention_residual, state))
4424}
4425
4426#[allow(clippy::too_many_arguments)]
4427fn execute_hyper_layer(
4428    layer: &memra_gguf::model_plan::LayerPlan,
4429    weights: &ReferenceWeights,
4430    input: &[f32],
4431    token_ids: &[u32],
4432    tokens: usize,
4433    hidden: usize,
4434    vocab: usize,
4435    streams: usize,
4436    epsilon: f32,
4437    sinkhorn_iterations: u32,
4438) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
4439    if input.len() != tokens * streams * hidden {
4440        return Err(ReferenceError::InvalidPlan {
4441            layer: Some(layer.index),
4442            reason: "HyperConnections input does not match tokens x streams x hidden",
4443        });
4444    }
4445    let attention_set = hyper_set(
4446        weights,
4447        layer.index,
4448        streams,
4449        hidden,
4450        LayerTensor::HyperAttentionFunction,
4451        LayerTensor::HyperAttentionBase,
4452        LayerTensor::HyperAttentionScale,
4453    )?;
4454    let (attention_input, post, combination) = memra_gguf::dsv4_forward::hc_pre(
4455        input,
4456        tokens,
4457        streams,
4458        hidden,
4459        &attention_set,
4460        sinkhorn_iterations,
4461        epsilon,
4462    );
4463    let attention_input = rms_norm(
4464        &attention_input,
4465        tokens,
4466        hidden,
4467        tensor(
4468            weights,
4469            &layer_id(layer.index, LayerTensor::PreAttentionNorm),
4470            &[hidden],
4471        )?,
4472        layer.pre_attention_norm.epsilon,
4473    );
4474    let (attention, state) = match &layer.attention {
4475        AttentionPlan::Full(attention) => full_attention(
4476            layer.index,
4477            attention,
4478            None,
4479            layer.pre_attention_norm.epsilon,
4480            weights,
4481            &attention_input,
4482            tokens,
4483            hidden,
4484            None,
4485        )?,
4486        AttentionPlan::SlidingWindow { attention, window } => full_attention(
4487            layer.index,
4488            attention,
4489            Some(*window as usize),
4490            layer.pre_attention_norm.epsilon,
4491            weights,
4492            &attention_input,
4493            tokens,
4494            hidden,
4495            None,
4496        )?,
4497        AttentionPlan::Mla(mla) => mla_attention(
4498            layer.index,
4499            mla,
4500            layer.pre_attention_norm.epsilon,
4501            weights,
4502            &attention_input,
4503            tokens,
4504            hidden,
4505        )?,
4506        AttentionPlan::GatedDeltaNet(gdn) => gated_delta_net(
4507            layer.index,
4508            gdn,
4509            layer.pre_attention_norm.epsilon,
4510            weights,
4511            &attention_input,
4512            tokens,
4513            hidden,
4514        )?,
4515        AttentionPlan::KimiDeltaNet(kda) => kimi_delta_net(
4516            layer.index,
4517            kda,
4518            layer.pre_attention_norm.epsilon,
4519            weights,
4520            &attention_input,
4521            tokens,
4522            hidden,
4523        )?,
4524    };
4525    let attention_residual = memra_gguf::dsv4_forward::hc_post(
4526        &attention,
4527        input,
4528        tokens,
4529        streams,
4530        hidden,
4531        &post,
4532        &combination,
4533    );
4534    if crate::hidden_trace::enabled() {
4535        let index = layer.index as i64;
4536        crate::hidden_trace::emit_last_row("mixer", index, tokens, hidden, &attention);
4537        crate::hidden_trace::emit_last_row(
4538            "attn",
4539            index,
4540            tokens,
4541            streams * hidden,
4542            &attention_residual,
4543        );
4544    }
4545
4546    let mlp_set = hyper_set(
4547        weights,
4548        layer.index,
4549        streams,
4550        hidden,
4551        LayerTensor::HyperMlpFunction,
4552        LayerTensor::HyperMlpBase,
4553        LayerTensor::HyperMlpScale,
4554    )?;
4555    let (mlp_input, post, combination) = memra_gguf::dsv4_forward::hc_pre(
4556        &attention_residual,
4557        tokens,
4558        streams,
4559        hidden,
4560        &mlp_set,
4561        sinkhorn_iterations,
4562        epsilon,
4563    );
4564    let mlp_input = rms_norm(
4565        &mlp_input,
4566        tokens,
4567        hidden,
4568        tensor(
4569            weights,
4570            &layer_id(layer.index, LayerTensor::PreMlpNorm),
4571            &[hidden],
4572        )?,
4573        layer.pre_mlp_norm.epsilon,
4574    );
4575    let mlp = match &layer.mlp {
4576        MlpPlan::Dense(mlp) => dense_mlp(layer.index, mlp, weights, &mlp_input, tokens, hidden)?,
4577        MlpPlan::Moe(moe) => moe_mlp(
4578            layer.index,
4579            moe,
4580            weights,
4581            &mlp_input,
4582            token_ids,
4583            tokens,
4584            hidden,
4585            vocab,
4586        )?,
4587    };
4588    let output = memra_gguf::dsv4_forward::hc_post(
4589        &mlp,
4590        &attention_residual,
4591        tokens,
4592        streams,
4593        hidden,
4594        &post,
4595        &combination,
4596    );
4597    if crate::hidden_trace::enabled() {
4598        let index = layer.index as i64;
4599        crate::hidden_trace::emit_last_row("ffn", index, tokens, hidden, &mlp);
4600        crate::hidden_trace::emit_last_row("layer", index, tokens, streams * hidden, &output);
4601    }
4602    Ok((output, state))
4603}
4604
4605#[allow(clippy::too_many_arguments)]
4606fn hyper_set(
4607    weights: &ReferenceWeights,
4608    layer: u32,
4609    streams: usize,
4610    hidden: usize,
4611    function: LayerTensor,
4612    base: LayerTensor,
4613    scale: LayerTensor,
4614) -> Result<memra_gguf::dsv4_forward::HcSet, ReferenceError> {
4615    let rows = (2 + streams) * streams;
4616    Ok(memra_gguf::dsv4_forward::HcSet {
4617        rows,
4618        fn_w: tensor(
4619            weights,
4620            &layer_id(layer, function),
4621            &[rows, streams * hidden],
4622        )?
4623        .to_vec(),
4624        base: tensor(weights, &layer_id(layer, base), &[rows])?.to_vec(),
4625        scale: tensor(weights, &layer_id(layer, scale), &[3])?.to_vec(),
4626    })
4627}
4628
4629/// One qwen4_exp gated-residual decoder layer (modular_qwen4_exp.py L796-833): optional
4630/// PLE add into the wide stream, attention read gate -> token mixer (QSA full attention
4631/// under the indexer mask, or GDN) -> per-stream write injection, then the same read /
4632/// mix / write around the MoE. There are NO input_layernorm modules in this family — the
4633/// read gate's grouped hc_norm IS the sublayer normalization.
4634#[allow(clippy::too_many_arguments)]
4635fn execute_gated_residual_layer(
4636    layer: &memra_gguf::model_plan::LayerPlan,
4637    weights: &ReferenceWeights,
4638    input: &[f32],
4639    token_ids: &[u32],
4640    tokens: usize,
4641    hidden: usize,
4642    vocab: usize,
4643    streams: usize,
4644    rank: usize,
4645    scope: LayerScope,
4646) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
4647    let wide = streams * hidden;
4648    if streams == 0 || rank == 0 || input.len() != tokens * wide {
4649        return Err(ReferenceError::InvalidPlan {
4650            layer: Some(layer.index),
4651            reason: "gated-residual input does not match tokens x streams x hidden",
4652        });
4653    }
4654    let prefix = scope.layer_prefix(layer.index);
4655    let epsilon = layer.pre_attention_norm.epsilon;
4656    let mut wide_state = input.to_vec();
4657    if let Some(ple) = layer.ple.as_ref() {
4658        // PLE adds to the wide stream BEFORE the attention read gate (modular L806-809);
4659        // the write gates below re-read the PLE-augmented stream as their hyper input.
4660        let ple_out = ple_block(
4661            layer.index,
4662            ple,
4663            epsilon,
4664            weights,
4665            &prefix,
4666            &wide_state,
4667            token_ids,
4668            tokens,
4669            streams,
4670            hidden,
4671        )?;
4672        add_in_place(&mut wide_state, &ple_out);
4673    }
4674    let (mixed, inject) = gated_residual_read(
4675        weights,
4676        &prefix,
4677        "attn_hyper_connection.",
4678        &wide_state,
4679        tokens,
4680        streams,
4681        hidden,
4682        rank,
4683        epsilon,
4684        true,
4685    )?;
4686    let (block_out, state) = match &layer.attention {
4687        AttentionPlan::Full(attention) => {
4688            let selection = layer
4689                .sparse_overlay
4690                .as_ref()
4691                .map(|overlay| {
4692                    micro_block_selection_mask(
4693                        layer.index,
4694                        overlay,
4695                        &attention.rope,
4696                        epsilon,
4697                        weights,
4698                        &prefix,
4699                        &mixed,
4700                        tokens,
4701                        hidden,
4702                    )
4703                })
4704                .transpose()?;
4705            full_attention(
4706                layer.index,
4707                attention,
4708                None,
4709                epsilon,
4710                weights,
4711                &mixed,
4712                tokens,
4713                hidden,
4714                selection.as_deref(),
4715            )?
4716        }
4717        AttentionPlan::GatedDeltaNet(gdn) => {
4718            gated_delta_net(layer.index, gdn, epsilon, weights, &mixed, tokens, hidden)?
4719        }
4720        _ => {
4721            return Err(ReferenceError::UnsupportedOperation {
4722                layer: Some(layer.index),
4723                operation: "gated-residual token mixer other than QSA/GDN",
4724            });
4725        }
4726    };
4727    gated_residual_write(
4728        &mut wide_state,
4729        &block_out,
4730        &inject,
4731        tokens,
4732        streams,
4733        hidden,
4734    );
4735    let (mixed, inject) = gated_residual_read(
4736        weights,
4737        &prefix,
4738        "mlp_hyper_connection.",
4739        &wide_state,
4740        tokens,
4741        streams,
4742        hidden,
4743        rank,
4744        layer.pre_mlp_norm.epsilon,
4745        true,
4746    )?;
4747    let mlp = match &layer.mlp {
4748        MlpPlan::Dense(mlp) => dense_mlp(layer.index, mlp, weights, &mixed, tokens, hidden)?,
4749        MlpPlan::Moe(moe) => moe_mlp(
4750            layer.index,
4751            moe,
4752            weights,
4753            &mixed,
4754            token_ids,
4755            tokens,
4756            hidden,
4757            vocab,
4758        )?,
4759    };
4760    gated_residual_write(&mut wide_state, &mlp, &inject, tokens, streams, hidden);
4761    Ok((wide_state, state))
4762}
4763
4764/// Qwen4ExpTextGatedResidual read gate (modular L541-558): grouped (1+w) RMSNorm of the
4765/// wide stream, `w = sigmoid(up(silu(down(normed) / streams)))`, `mixed = mean over
4766/// streams of (w * normed)`, and — when `with_inject` — the write-injection scalars
4767/// `2 * sigmoid(block_inject(normed) / streams)` from the SAME normed input. The exit /
4768/// MTP mixer is the same read with `with_inject = false` (use_combine=False). Returned
4769/// inject is `[tokens, streams]` (empty when not requested).
4770#[allow(clippy::too_many_arguments)]
4771fn gated_residual_read(
4772    weights: &ReferenceWeights,
4773    prefix: &str,
4774    sublayer: &str,
4775    x: &[f32],
4776    tokens: usize,
4777    streams: usize,
4778    hidden: usize,
4779    rank: usize,
4780    epsilon: f32,
4781    with_inject: bool,
4782) -> Result<(Vec<f32>, Vec<f32>), ReferenceError> {
4783    let wide = streams * hidden;
4784    if x.len() != tokens * wide || streams == 0 || rank == 0 {
4785        return Err(ReferenceError::InvalidPlan {
4786            layer: None,
4787            reason: "gated-residual read requires tokens x streams x hidden input",
4788        });
4789    }
4790    let norm = tensor(
4791        weights,
4792        &qwen4exp_family_id(format!("{prefix}{sublayer}hc_norm.weight")),
4793        &[wide],
4794    )?;
4795    let down = tensor(
4796        weights,
4797        &qwen4exp_family_id(format!("{prefix}{sublayer}input_mix_weight_down.weight")),
4798        &[rank, wide],
4799    )?;
4800    let up = tensor(
4801        weights,
4802        &qwen4exp_family_id(format!("{prefix}{sublayer}input_mix_weight_up.weight")),
4803        &[wide, rank],
4804    )?;
4805    let inject_weight = with_inject
4806        .then(|| {
4807            tensor(
4808                weights,
4809                &qwen4exp_family_id(format!("{prefix}{sublayer}block_inject_weight.weight")),
4810                &[streams, wide],
4811            )
4812        })
4813        .transpose()?;
4814    let normed = grouped_rms_norm(x, tokens, streams, hidden, norm, epsilon);
4815    let mut mixed = vec![0.0; tokens * hidden];
4816    let mut inject = vec![0.0; if with_inject { tokens * streams } else { 0 }];
4817    for token in 0..tokens {
4818        let row = &normed[token * wide..(token + 1) * wide];
4819        let mut low = vec![0.0; rank];
4820        for index in 0..rank {
4821            let mut sum = 0.0;
4822            for dim in 0..wide {
4823                sum += down[index * wide + dim] * row[dim];
4824            }
4825            low[index] = silu(sum / streams as f32);
4826        }
4827        for column in 0..hidden {
4828            let mut sum = 0.0;
4829            for stream in 0..streams {
4830                let dim = stream * hidden + column;
4831                let mut gate = 0.0;
4832                for index in 0..rank {
4833                    gate += up[dim * rank + index] * low[index];
4834                }
4835                sum += sigmoid(gate) * row[dim];
4836            }
4837            mixed[token * hidden + column] = sum / streams as f32;
4838        }
4839        if let Some(inject_weight) = inject_weight {
4840            for stream in 0..streams {
4841                let mut sum = 0.0;
4842                for dim in 0..wide {
4843                    sum += inject_weight[stream * wide + dim] * row[dim];
4844                }
4845                inject[token * streams + stream] = 2.0 * sigmoid(sum / streams as f32);
4846            }
4847        }
4848    }
4849    Ok((mixed, inject))
4850}
4851
4852/// Write half of the gated residual (modular L825-826): the wide stream gains the outer
4853/// product `block_out ⊗ inject` — stream s receives `block_out * inject[s]` — on top of
4854/// the PRE-norm hyper input.
4855fn gated_residual_write(
4856    wide_state: &mut [f32],
4857    block_out: &[f32],
4858    inject: &[f32],
4859    tokens: usize,
4860    streams: usize,
4861    hidden: usize,
4862) {
4863    for token in 0..tokens {
4864        for stream in 0..streams {
4865            let weight = inject[token * streams + stream];
4866            let offset = token * streams * hidden + stream * hidden;
4867            for column in 0..hidden {
4868                wide_state[offset + column] += block_out[token * hidden + column] * weight;
4869            }
4870        }
4871    }
4872}
4873
4874/// Qwen4ExpTextRMSNorm with group_size = hidden (modular L298-309): every stream group of
4875/// the wide vector normalizes independently; `weight` spans the FULL wide width. Weights
4876/// are EFFECTIVE (1+w) — the checkpoint ships zero-centered values (the family convention,
4877/// modular L859-861 zero-init receipt) folded at binding, like every norm in this crate.
4878fn grouped_rms_norm(
4879    x: &[f32],
4880    tokens: usize,
4881    streams: usize,
4882    hidden: usize,
4883    weight: &[f32],
4884    epsilon: f32,
4885) -> Vec<f32> {
4886    let wide = streams * hidden;
4887    let mut result = vec![0.0; x.len()];
4888    for token in 0..tokens {
4889        for stream in 0..streams {
4890            let offset = token * wide + stream * hidden;
4891            let group = &x[offset..offset + hidden];
4892            let mean_square = group.iter().map(|value| value * value).sum::<f32>() / hidden as f32;
4893            let inverse = 1.0 / (mean_square + epsilon).sqrt();
4894            for column in 0..hidden {
4895                result[offset + column] =
4896                    group[column] * inverse * weight[stream * hidden + column];
4897            }
4898        }
4899    }
4900    result
4901}
4902
4903/// QSA indexer selection (modular L367-473): the fused `index_qk_proj` splits into
4904/// per-head queries (per-head RMSNorm, then the MAIN partial rope at the query position)
4905/// and ONE shared RAW key per token (cached pre-norm, pre-rope). Per query token the
4906/// visible tokens (causal here — the reference sees the whole prompt) form complete
4907/// blocks of `block_size`; each block pools its raw keys by fp32 mean -> k_layernorm ->
4908/// rope at the block's FIRST position; `score = Σ_heads relu(q·k) / sqrt(head_dim)`; the
4909/// top `min(budget_blocks, complete)` blocks stay visible plus the always-visible
4910/// incomplete tail.
4911///
4912/// Tie rule — DELIBERATE PIN: score descending, then block index ascending. torch.topk's
4913/// tie order is implementation-defined (SEMANTICS.md §QSA indexer), so the reference pins
4914/// a total order; parity fixtures must be tie-free (dsv4-lane lesson).
4915#[allow(clippy::too_many_arguments)]
4916fn micro_block_selection_mask(
4917    layer: u32,
4918    overlay: &MicroBlockIndexPlan,
4919    rope: &RopePlan,
4920    epsilon: f32,
4921    weights: &ReferenceWeights,
4922    prefix: &str,
4923    x: &[f32],
4924    tokens: usize,
4925    hidden: usize,
4926) -> Result<Vec<bool>, ReferenceError> {
4927    let heads = overlay.query_heads as usize;
4928    let kv_heads = overlay.kv_heads as usize;
4929    let head_dim = overlay.head_dim as usize;
4930    let block_size = overlay.block_size as usize;
4931    let budget_blocks = overlay.budget_blocks as usize;
4932    if heads == 0 || head_dim == 0 || block_size == 0 || budget_blocks == 0 {
4933        return Err(ReferenceError::InvalidPlan {
4934            layer: Some(layer),
4935            reason: "micro-block indexer requires heads, head_dim, block size, and budget",
4936        });
4937    }
4938    if kv_heads != 1 {
4939        // modular L406 squeezes exactly one shared key head; more is a different program.
4940        return Err(ReferenceError::UnsupportedOperation {
4941            layer: Some(layer),
4942            operation: "micro-block indexer with more than one key head",
4943        });
4944    }
4945    let qk_width = (heads + kv_heads) * head_dim;
4946    let projected = linear(
4947        x,
4948        tensor(
4949            weights,
4950            &qwen4exp_family_id(format!("{prefix}self_attn.indexer.index_qk_proj.weight")),
4951            &[qk_width, hidden],
4952        )?,
4953        tokens,
4954        hidden,
4955        qk_width,
4956    );
4957    let q_norm_weight = tensor(
4958        weights,
4959        &qwen4exp_family_id(format!("{prefix}self_attn.indexer.q_layernorm.weight")),
4960        &[head_dim],
4961    )?;
4962    let k_norm_weight = tensor(
4963        weights,
4964        &qwen4exp_family_id(format!("{prefix}self_attn.indexer.k_layernorm.weight")),
4965        &[head_dim],
4966    )?;
4967    let mut query = vec![0.0; tokens * heads * head_dim];
4968    let mut raw_keys = vec![0.0; tokens * head_dim];
4969    for token in 0..tokens {
4970        query[token * heads * head_dim..(token + 1) * heads * head_dim]
4971            .copy_from_slice(&projected[token * qk_width..token * qk_width + heads * head_dim]);
4972        raw_keys[token * head_dim..(token + 1) * head_dim].copy_from_slice(
4973            &projected[token * qk_width + heads * head_dim..(token + 1) * qk_width],
4974        );
4975    }
4976    let mut query = rms_norm(&query, tokens * heads, head_dim, q_norm_weight, epsilon);
4977    // The indexer consumes the MAIN rotary cos/sin, partial over rope_dimensions of the
4978    // (wider) index head; text-only mrope degenerates to plain partial rope.
4979    let rope_dims = overlay.rope_dimensions as usize;
4980    let (factors, mscale) = rope_factor_values(rope, weights)?;
4981    apply_rope(
4982        &mut query,
4983        tokens,
4984        heads,
4985        head_dim,
4986        rope_dims,
4987        rope.base,
4988        factors.as_deref(),
4989        mscale,
4990    );
4991
4992    let mut mask = vec![false; tokens * tokens];
4993    let scale = (head_dim as f32).sqrt();
4994    for token in 0..tokens {
4995        let visible = token + 1;
4996        let complete = visible / block_size;
4997        let mut scored: Vec<(usize, f32)> = Vec::with_capacity(complete);
4998        for block in 0..complete {
4999            let start = block * block_size;
5000            // fp32 mean of the RAW keys (modular L437), then k_layernorm, then rope at
5001            // the block-start position (group_starts, L439-444).
5002            let mut pooled = vec![0.0f32; head_dim];
5003            for offset in 0..block_size {
5004                for dim in 0..head_dim {
5005                    pooled[dim] += raw_keys[(start + offset) * head_dim + dim];
5006                }
5007            }
5008            for value in &mut pooled {
5009                *value /= block_size as f32;
5010            }
5011            let mut pooled = rms_norm(&pooled, 1, head_dim, k_norm_weight, epsilon);
5012            apply_rope_at_position(
5013                &mut pooled,
5014                1,
5015                head_dim,
5016                rope_dims,
5017                rope.base,
5018                factors.as_deref(),
5019                mscale,
5020                start,
5021            );
5022            let mut score = 0.0f32;
5023            for head in 0..heads {
5024                let mut dot = 0.0f32;
5025                for dim in 0..head_dim {
5026                    dot += query[(token * heads + head) * head_dim + dim] * pooled[dim];
5027                }
5028                score += dot.max(0.0);
5029            }
5030            scored.push((block, score / scale));
5031        }
5032        scored.sort_by(|left, right| right.1.total_cmp(&left.1).then(left.0.cmp(&right.0)));
5033        for &(block, _) in scored.iter().take(budget_blocks.min(complete)) {
5034            for offset in 0..block_size {
5035                mask[token * tokens + block * block_size + offset] = true;
5036            }
5037        }
5038        // The incomplete tail block is always selected (modular L456-457) — this is also
5039        // what guarantees every query keeps at least one visible source when its own
5040        // block is complete but unselected... except at exact block boundaries, where
5041        // topk >= 1 block always fires (complete >= 1).
5042        for source in complete * block_size..visible {
5043            mask[token * tokens + source] = true;
5044        }
5045    }
5046    Ok(mask)
5047}
5048
5049/// qwen4_exp PLE block (modular L706-778): gather the hashed n-gram rows, key them
5050/// against the wide stream per-stream (signed-sqrt sigmoid gates), then add a dilated
5051/// depthwise causal conv refinement. The reference processes the whole prompt in one
5052/// pass; the token-history semantics stay exact (the first max_ngram-1 context positions
5053/// read EOS).
5054#[allow(clippy::too_many_arguments)]
5055fn ple_block(
5056    layer: u32,
5057    plan: &PleEmbeddingPlan,
5058    epsilon: f32,
5059    weights: &ReferenceWeights,
5060    prefix: &str,
5061    wide_state: &[f32],
5062    token_ids: &[u32],
5063    tokens: usize,
5064    streams: usize,
5065    hidden: usize,
5066) -> Result<Vec<f32>, ReferenceError> {
5067    let heads = plan.ngram_heads as usize;
5068    let head_dim = plan.head_embed_dim as usize;
5069    let embed_dim = plan.embed_dim as usize;
5070    let kernel = plan.conv_kernel as usize;
5071    let max_ngram = plan.max_ngram as usize;
5072    let wide = streams * hidden;
5073    if heads == 0
5074        || head_dim == 0
5075        || kernel == 0
5076        || max_ngram < 2
5077        || embed_dim != heads * head_dim
5078        || !heads.is_multiple_of(max_ngram - 1)
5079    {
5080        return Err(ReferenceError::InvalidPlan {
5081            layer: Some(layer),
5082            reason: "PLE requires consistent n-gram head geometry",
5083        });
5084    }
5085    let multipliers = tensor_i64(
5086        weights,
5087        &qwen4exp_family_id(format!("{prefix}ple.ple_embedding.layer_multipliers")),
5088        &[max_ngram],
5089    )?;
5090    let sizes = tensor_i64(
5091        weights,
5092        &qwen4exp_family_id(format!("{prefix}ple.ple_embedding.ngram_heads_vocab_sizes")),
5093        &[heads],
5094    )?;
5095    let offsets = tensor_i64(
5096        weights,
5097        &qwen4exp_family_id(format!("{prefix}ple.ple_embedding.ngram_heads_offsets")),
5098        &[heads],
5099    )?;
5100    let ids = ngram_ids(
5101        token_ids,
5102        multipliers,
5103        sizes,
5104        offsets,
5105        max_ngram,
5106        heads / (max_ngram - 1),
5107        plan.eos_token_id,
5108        layer,
5109    )?;
5110    let table_id = qwen4exp_family_id(format!("{prefix}ple.ple_embedding.ngram_embedding"));
5111    let table = weights
5112        .get(&table_id)
5113        .ok_or_else(|| ReferenceError::MissingTensor(table_id.clone()))?;
5114    let rows = table.shape.first().copied().unwrap_or(0);
5115    if table.shape.len() != 2 || table.shape[1] != head_dim || table.data.len() != rows * head_dim {
5116        return Err(ReferenceError::TensorShape {
5117            id: Some(table_id),
5118            expected: vec![rows, head_dim],
5119            actual_elements: table.data.len(),
5120        });
5121    }
5122    let mut embeddings = vec![0.0; tokens * embed_dim];
5123    for token in 0..tokens {
5124        for head in 0..heads {
5125            let id = ids[token * heads + head];
5126            if id < 0 || id as usize >= rows {
5127                return Err(ReferenceError::InvalidPlan {
5128                    layer: Some(layer),
5129                    reason: "n-gram id addressed outside the embedding table",
5130                });
5131            }
5132            let target = token * embed_dim + head * head_dim;
5133            embeddings[target..target + head_dim]
5134                .copy_from_slice(&table.data[id as usize * head_dim..(id as usize + 1) * head_dim]);
5135        }
5136    }
5137    let key = linear(
5138        &embeddings,
5139        tensor(
5140            weights,
5141            &qwen4exp_family_id(format!("{prefix}ple.key_proj.weight")),
5142            &[wide, embed_dim],
5143        )?,
5144        tokens,
5145        embed_dim,
5146        wide,
5147    );
5148    let key = grouped_rms_norm(
5149        &key,
5150        tokens,
5151        streams,
5152        hidden,
5153        tensor(
5154            weights,
5155            &qwen4exp_family_id(format!("{prefix}ple.norm_key.weight")),
5156            &[wide],
5157        )?,
5158        epsilon,
5159    );
5160    let value = linear(
5161        &embeddings,
5162        tensor(
5163            weights,
5164            &qwen4exp_family_id(format!("{prefix}ple.value_proj.weight")),
5165            &[hidden, embed_dim],
5166        )?,
5167        tokens,
5168        embed_dim,
5169        hidden,
5170    );
5171    let query = grouped_rms_norm(
5172        wide_state,
5173        tokens,
5174        streams,
5175        hidden,
5176        tensor(
5177            weights,
5178            &qwen4exp_family_id(format!("{prefix}ple.norm_query.weight")),
5179            &[wide],
5180        )?,
5181        epsilon,
5182    );
5183    let mut gated_value = vec![0.0; tokens * wide];
5184    for token in 0..tokens {
5185        for stream in 0..streams {
5186            let offset = token * wide + stream * hidden;
5187            let mut dot = 0.0;
5188            for column in 0..hidden {
5189                dot += key[offset + column] * query[offset + column];
5190            }
5191            let gate = dot / (hidden as f32).sqrt();
5192            // signed sqrt (modular L770): sqrt(clamp_min(|g|, 1e-6)) * sign(g); torch
5193            // sign(0) = 0, so a zero gate stays zero (f32::signum would say +1).
5194            let magnitude = gate.abs().max(1e-6).sqrt();
5195            let gate = if gate > 0.0 {
5196                magnitude
5197            } else if gate < 0.0 {
5198                -magnitude
5199            } else {
5200                0.0
5201            };
5202            let gate = sigmoid(gate);
5203            for column in 0..hidden {
5204                gated_value[offset + column] = gate * value[token * hidden + column];
5205            }
5206        }
5207    }
5208    let normed = grouped_rms_norm(
5209        &gated_value,
5210        tokens,
5211        streams,
5212        hidden,
5213        tensor(
5214            weights,
5215            &qwen4exp_family_id(format!("{prefix}ple.norm_conv.weight")),
5216            &[wide],
5217        )?,
5218        epsilon,
5219    );
5220    // Depthwise causal conv over the NORMED gated value: kernel taps sit `dilation`
5221    // (= max_ngram) apart, left-pad (kernel-1)*dilation (modular L739-756: conv weight
5222    // [wide, 1, kernel] — consumed squeezed like the GDN conv row).
5223    let conv_weight = tensor(
5224        weights,
5225        &qwen4exp_family_id(format!("{prefix}ple.conv1d.weight")),
5226        &[wide, kernel],
5227    )?;
5228    let dilation = max_ngram;
5229    let mut output = gated_value;
5230    for token in 0..tokens {
5231        for channel in 0..wide {
5232            let mut sum = 0.0;
5233            for tap in 0..kernel {
5234                let reach = ((kernel - 1 - tap) * dilation) as isize;
5235                let source = token as isize - reach;
5236                if source >= 0 {
5237                    sum += normed[source as usize * wide + channel]
5238                        * conv_weight[channel * kernel + tap];
5239                }
5240            }
5241            output[token * wide + channel] += silu(sum);
5242        }
5243    }
5244    Ok(output)
5245}
5246
5247/// N-gram ids (modular L642-703): token history = (max_ngram-1) EOS context positions ++
5248/// prompt; `shifted[j]` shifts right by j with EOS-segment reset; for n in 2..=max_ngram
5249/// the shifted ids mix by wrapping-i64 multiply + XOR, and each of that n-gram's heads
5250/// takes `mixed mod head_vocab_size + head_offset` (torch.remainder = floor mod; the
5251/// divisors are positive so rem_euclid matches). Multipliers / sizes / offsets are
5252/// checkpoint I64 buffers — LOADED, never re-derived (SEMANTICS.md §PLE). Returns
5253/// `[tokens, total_heads]` (history context rows dropped).
5254#[allow(clippy::too_many_arguments)]
5255fn ngram_ids(
5256    token_ids: &[u32],
5257    multipliers: &[i64],
5258    sizes: &[i64],
5259    offsets: &[i64],
5260    max_ngram: usize,
5261    heads_per_ngram: usize,
5262    eos_token_id: u32,
5263    layer: u32,
5264) -> Result<Vec<i64>, ReferenceError> {
5265    let context = max_ngram - 1;
5266    let eos = eos_token_id as i64;
5267    let total_heads = (max_ngram - 1) * heads_per_ngram;
5268    if multipliers.len() != max_ngram || sizes.len() != total_heads || offsets.len() != total_heads
5269    {
5270        return Err(ReferenceError::InvalidPlan {
5271            layer: Some(layer),
5272            reason: "n-gram index buffers do not match the head geometry",
5273        });
5274    }
5275    if sizes.iter().any(|&size| size <= 0) || offsets.iter().any(|&offset| offset < 0) {
5276        return Err(ReferenceError::InvalidPlan {
5277            layer: Some(layer),
5278            reason: "n-gram head vocab sizes must be positive and offsets non-negative",
5279        });
5280    }
5281    let mut history = Vec::with_capacity(context + token_ids.len());
5282    history.extend(std::iter::repeat_n(eos, context));
5283    history.extend(token_ids.iter().map(|&token| token as i64));
5284    let shifted: Vec<Vec<i64>> = (0..max_ngram)
5285        .map(|shift| shift_right_ignore_eos(&history, shift, eos))
5286        .collect();
5287    let tokens = token_ids.len();
5288    let mut ids = vec![0i64; tokens * total_heads];
5289    for ngram in 2..=max_ngram {
5290        let head_start = (ngram - 2) * heads_per_ngram;
5291        for token in 0..tokens {
5292            let position = context + token;
5293            let mut mixed = shifted[0][position].wrapping_mul(multipliers[0]);
5294            for shift in 1..ngram {
5295                mixed ^= shifted[shift][position].wrapping_mul(multipliers[shift]);
5296            }
5297            for head in 0..heads_per_ngram {
5298                let index = head_start + head;
5299                ids[token * total_heads + index] = mixed.rem_euclid(sizes[index]) + offsets[index];
5300            }
5301        }
5302    }
5303    Ok(ids)
5304}
5305
5306/// modular L642-656: positions whose in-segment index (counted from the token after the
5307/// EOS strictly before them) is smaller than the shift — or whose shifted source
5308/// underflows the history — read EOS instead of a cross-segment token.
5309fn shift_right_ignore_eos(history: &[i64], shift: usize, eos: i64) -> Vec<i64> {
5310    if shift == 0 {
5311        return history.to_vec();
5312    }
5313    let mut last_eos_inclusive: i64 = -1;
5314    let mut output = Vec::with_capacity(history.len());
5315    for (position, &token) in history.iter().enumerate() {
5316        let previous_eos = last_eos_inclusive;
5317        if token == eos {
5318            last_eos_inclusive = position as i64;
5319        }
5320        let segment_start = previous_eos + 1;
5321        let position_in_segment = position as i64 - segment_start;
5322        let source = position as i64 - shift as i64;
5323        let valid = position_in_segment >= shift as i64 && source >= 0;
5324        output.push(if valid { history[source as usize] } else { eos });
5325    }
5326    output
5327}
5328
5329fn tensor_i64<'a>(
5330    weights: &'a ReferenceWeights,
5331    id: &TensorId,
5332    expected: &[usize],
5333) -> Result<&'a [i64], ReferenceError> {
5334    let tensor = weights
5335        .get(id)
5336        .ok_or_else(|| ReferenceError::MissingTensor(id.clone()))?;
5337    let Some(ints) = tensor.ints.as_ref() else {
5338        return Err(ReferenceError::IntegerTensorRequired(id.clone()));
5339    };
5340    if tensor.shape != expected {
5341        return Err(ReferenceError::TensorShape {
5342            id: Some(id.clone()),
5343            expected: expected.to_vec(),
5344            actual_elements: ints.len(),
5345        });
5346    }
5347    Ok(ints)
5348}
5349
5350fn execute_gemma_dense_layer(
5351    layer: &memra_gguf::model_plan::LayerPlan,
5352    weights: &ReferenceWeights,
5353    input: &[f32],
5354    tokens: usize,
5355    hidden: usize,
5356) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
5357    let ResidualTopology::Gemma {
5358        post_attention_norm,
5359        post_mlp_norm,
5360        layer_scale,
5361        parallel_moe: None,
5362    } = layer.residual
5363    else {
5364        return Err(ReferenceError::UnsupportedOperation {
5365            layer: Some(layer.index),
5366            operation: "gemma parallel MoE residual",
5367        });
5368    };
5369    let pre_attn = rms_norm(
5370        input,
5371        tokens,
5372        hidden,
5373        tensor(
5374            weights,
5375            &layer_id(layer.index, LayerTensor::PreAttentionNorm),
5376            &[hidden],
5377        )?,
5378        layer.pre_attention_norm.epsilon,
5379    );
5380    let (attention, state) = match &layer.attention {
5381        AttentionPlan::Full(attention) => full_attention(
5382            layer.index,
5383            attention,
5384            None,
5385            layer.pre_attention_norm.epsilon,
5386            weights,
5387            &pre_attn,
5388            tokens,
5389            hidden,
5390            None,
5391        )?,
5392        AttentionPlan::SlidingWindow { attention, window } => full_attention(
5393            layer.index,
5394            attention,
5395            Some(*window as usize),
5396            layer.pre_attention_norm.epsilon,
5397            weights,
5398            &pre_attn,
5399            tokens,
5400            hidden,
5401            None,
5402        )?,
5403        _ => {
5404            return Err(ReferenceError::UnsupportedOperation {
5405                layer: Some(layer.index),
5406                operation: "gemma non-softmax attention",
5407            });
5408        }
5409    };
5410    let post_attention = rms_norm(
5411        &attention,
5412        tokens,
5413        hidden,
5414        tensor(
5415            weights,
5416            &layer_id(layer.index, LayerTensor::PostAttentionNorm),
5417            &[hidden],
5418        )?,
5419        post_attention_norm.epsilon,
5420    );
5421    let mut attention_residual = input.to_vec();
5422    add_in_place(&mut attention_residual, &post_attention);
5423    let pre_mlp = rms_norm(
5424        &attention_residual,
5425        tokens,
5426        hidden,
5427        tensor(
5428            weights,
5429            &layer_id(layer.index, LayerTensor::PreMlpNorm),
5430            &[hidden],
5431        )?,
5432        layer.pre_mlp_norm.epsilon,
5433    );
5434    let MlpPlan::Dense(mlp) = &layer.mlp else {
5435        return Err(ReferenceError::UnsupportedOperation {
5436            layer: Some(layer.index),
5437            operation: "gemma parallel MoE residual",
5438        });
5439    };
5440    let mlp = dense_mlp(layer.index, mlp, weights, &pre_mlp, tokens, hidden)?;
5441    let mlp = rms_norm(
5442        &mlp,
5443        tokens,
5444        hidden,
5445        tensor(
5446            weights,
5447            &layer_id(layer.index, LayerTensor::PostMlpNorm),
5448            &[hidden],
5449        )?,
5450        post_mlp_norm.epsilon,
5451    );
5452    let scale = match layer_scale {
5453        GemmaLayerScale::Learned => tensor(
5454            weights,
5455            &layer_id(layer.index, LayerTensor::LayerScale),
5456            &[1],
5457        )?[0],
5458    };
5459    let mut output = attention_residual;
5460    add_in_place(&mut output, &mlp);
5461    for value in &mut output {
5462        *value *= scale;
5463    }
5464    Ok((output, state))
5465}
5466
5467#[allow(clippy::too_many_arguments)]
5468/// Execute ONLY the MTP draft arm on CALLER-PROVIDED trunk wide states — the
5469/// real-checkpoint draft-parity instrument (mtp-spec lane): the full-trunk host
5470/// reference cannot hold the 360 GB artifact, but the MTP block's rows fit host f32,
5471/// so the engine's captured trunk wide state feeds this host twin and the draft
5472/// programs are compared row for row. `trunk_hidden` is [tokens, streams*hidden]
5473/// (gated-residual plans) or [tokens, hidden]; row i seeds token_ids[i] — the same
5474/// pairing `execute` uses internally.
5475pub fn execute_mtp_standalone(
5476    plan: &ModelPlan,
5477    weights: &ReferenceWeights,
5478    token_ids: &[u32],
5479    trunk_hidden: &[f32],
5480) -> Result<Vec<ReferenceMtpOutput>, ReferenceError> {
5481    let hidden = plan.hidden_size as usize;
5482    let vocab = plan.vocab_size as usize;
5483    let tokens = token_ids.len();
5484    let embedding = tensor(weights, &TensorId::TokenEmbedding, &[vocab, hidden])?;
5485    let output = weights
5486        .get(&TensorId::OutputProjection)
5487        .map(|tensor| tensor_checked(&TensorId::OutputProjection, tensor, &[vocab, hidden]))
5488        .transpose()?
5489        .unwrap_or(embedding);
5490    execute_mtp(
5491        plan,
5492        weights,
5493        token_ids,
5494        embedding,
5495        trunk_hidden,
5496        tokens,
5497        hidden,
5498        vocab,
5499        output,
5500    )
5501}
5502
5503// too_many_arguments: the MTP executor takes exactly the seams the trunk hands it;
5504// bundling them into a struct would reshape the oracle's call surface for a lint.
5505#[allow(clippy::too_many_arguments)]
5506fn execute_mtp(
5507    plan: &ModelPlan,
5508    weights: &ReferenceWeights,
5509    token_ids: &[u32],
5510    embedding: &[f32],
5511    trunk_hidden: &[f32],
5512    tokens: usize,
5513    hidden: usize,
5514    vocab: usize,
5515    model_output: &[f32],
5516) -> Result<Vec<ReferenceMtpOutput>, ReferenceError> {
5517    if plan.mtp_blocks.is_empty() {
5518        return Ok(Vec::new());
5519    }
5520    let gated = gated_residual_topology(plan)?;
5521    if gated.is_some() && plan.mtp_blocks.len() > 1 {
5522        // The checkpoint has ONE mtp.* namespace (glue + mixer); a second depth would
5523        // alias its tensors.
5524        return Err(ReferenceError::UnsupportedOperation {
5525            layer: None,
5526            operation: "multi-depth gated-residual MTP",
5527        });
5528    }
5529    let mut embedded = vec![0.0; tokens * hidden];
5530    for (position, &token) in token_ids.iter().enumerate() {
5531        let token = token as usize;
5532        embedded[position * hidden..(position + 1) * hidden]
5533            .copy_from_slice(&embedding[token * hidden..(token + 1) * hidden]);
5534    }
5535    let mut source_hidden = trunk_hidden.to_vec();
5536    let mut outputs = Vec::with_capacity(plan.mtp_blocks.len());
5537    for block in &plan.mtp_blocks {
5538        let fused = match block.input.fusion {
5539            memra_gguf::model_plan::MtpFusionPlan::ConcatenateProjection => {
5540                if source_hidden.len() != tokens * hidden {
5541                    return Err(ReferenceError::UnsupportedOperation {
5542                        layer: None,
5543                        operation: "HyperConnections MTP fusion",
5544                    });
5545                }
5546                let embedding_norm = rms_norm(
5547                    &embedded,
5548                    tokens,
5549                    hidden,
5550                    tensor(
5551                        weights,
5552                        &TensorId::Mtp {
5553                            depth: block.depth,
5554                            tensor: MtpTensor::EmbeddingNorm,
5555                        },
5556                        &[hidden],
5557                    )?,
5558                    block.input.embedding_norm.epsilon,
5559                );
5560                let hidden_norm = rms_norm(
5561                    &source_hidden,
5562                    tokens,
5563                    hidden,
5564                    tensor(
5565                        weights,
5566                        &TensorId::Mtp {
5567                            depth: block.depth,
5568                            tensor: MtpTensor::HiddenNorm,
5569                        },
5570                        &[hidden],
5571                    )?,
5572                    block.input.hidden_norm.epsilon,
5573                );
5574                let mut concatenated = vec![0.0; tokens * 2 * hidden];
5575                for token in 0..tokens {
5576                    concatenated[token * 2 * hidden..token * 2 * hidden + hidden]
5577                        .copy_from_slice(&embedding_norm[token * hidden..(token + 1) * hidden]);
5578                    concatenated[token * 2 * hidden + hidden..(token + 1) * 2 * hidden]
5579                        .copy_from_slice(&hidden_norm[token * hidden..(token + 1) * hidden]);
5580                }
5581                linear(
5582                    &concatenated,
5583                    tensor(
5584                        weights,
5585                        &TensorId::Mtp {
5586                            depth: block.depth,
5587                            tensor: MtpTensor::FusionProjection,
5588                        },
5589                        &[hidden, 2 * hidden],
5590                    )?,
5591                    tokens,
5592                    2 * hidden,
5593                    hidden,
5594                )
5595            }
5596            memra_gguf::model_plan::MtpFusionPlan::SeparateProjections => {
5597                // qwen4_exp (SEMANTICS.md §MTP, sglang_qwen4_exp_mtp.py L105-115): the
5598                // draft input is the trunk's WIDE state, normed FLAT over the full wide
5599                // vector (GemmaRMSNorm(hc_count*hidden) — not grouped), viewed per stream
5600                // through fc_hidden, plus fc_embedding(norm(embed)) broadcast over streams.
5601                let Some((streams, _)) = gated else {
5602                    return Err(ReferenceError::InvalidPlan {
5603                        layer: Some(block.layer.index),
5604                        reason: "separate-projection MTP fusion requires a gated-residual trunk",
5605                    });
5606                };
5607                let wide = streams * hidden;
5608                if source_hidden.len() != tokens * wide {
5609                    return Err(ReferenceError::InvalidPlan {
5610                        layer: Some(block.layer.index),
5611                        reason: "separate-projection MTP fusion requires the wide trunk state",
5612                    });
5613                }
5614                let embedding_norm = rms_norm(
5615                    &embedded,
5616                    tokens,
5617                    hidden,
5618                    tensor(
5619                        weights,
5620                        &TensorId::Mtp {
5621                            depth: block.depth,
5622                            tensor: MtpTensor::EmbeddingNorm,
5623                        },
5624                        &[hidden],
5625                    )?,
5626                    block.input.embedding_norm.epsilon,
5627                );
5628                let embedding_projected = linear(
5629                    &embedding_norm,
5630                    tensor(
5631                        weights,
5632                        &TensorId::Mtp {
5633                            depth: block.depth,
5634                            tensor: MtpTensor::EmbeddingProjection,
5635                        },
5636                        &[hidden, hidden],
5637                    )?,
5638                    tokens,
5639                    hidden,
5640                    hidden,
5641                );
5642                let hidden_norm = rms_norm(
5643                    &source_hidden,
5644                    tokens,
5645                    wide,
5646                    tensor(
5647                        weights,
5648                        &TensorId::Mtp {
5649                            depth: block.depth,
5650                            tensor: MtpTensor::HiddenNorm,
5651                        },
5652                        &[wide],
5653                    )?,
5654                    block.input.hidden_norm.epsilon,
5655                );
5656                let hidden_projected = linear(
5657                    &hidden_norm,
5658                    tensor(
5659                        weights,
5660                        &TensorId::Mtp {
5661                            depth: block.depth,
5662                            tensor: MtpTensor::HiddenProjection,
5663                        },
5664                        &[hidden, hidden],
5665                    )?,
5666                    tokens * streams,
5667                    hidden,
5668                    hidden,
5669                );
5670                let mut fused = hidden_projected;
5671                for token in 0..tokens {
5672                    for stream in 0..streams {
5673                        for column in 0..hidden {
5674                            fused[(token * streams + stream) * hidden + column] +=
5675                                embedding_projected[token * hidden + column];
5676                        }
5677                    }
5678                }
5679                fused
5680            }
5681        };
5682        let (hidden_next, state) = execute_layer(
5683            &block.layer,
5684            weights,
5685            &fused,
5686            token_ids,
5687            tokens,
5688            hidden,
5689            vocab,
5690            LayerScope::Mtp { depth: block.depth },
5691        )?;
5692        let norm_id = TensorId::Mtp {
5693            depth: block.depth,
5694            tensor: MtpTensor::OutputNorm,
5695        };
5696        let final_hidden = if let Some((streams, rank)) = gated {
5697            // The draft exits through its OWN hyper_connection_mixer (SEMANTICS.md §MTP);
5698            // there is no MTP final norm and no model OutputNorm to fall back to.
5699            gated_residual_read(
5700                weights,
5701                LayerScope::Mtp { depth: block.depth }.mixer_prefix(),
5702                "",
5703                &hidden_next,
5704                tokens,
5705                streams,
5706                hidden,
5707                rank,
5708                plan.output_norm.epsilon,
5709                false,
5710            )?
5711            .0
5712        } else {
5713            let norm = match weights.get(&norm_id) {
5714                Some(tensor) => tensor_checked(&norm_id, tensor, &[hidden])?,
5715                None => tensor(weights, &TensorId::OutputNorm, &[hidden])?,
5716            };
5717            rms_norm(&hidden_next, tokens, hidden, norm, plan.output_norm.epsilon)
5718        };
5719        let head_id = TensorId::Mtp {
5720            depth: block.depth,
5721            tensor: MtpTensor::OutputProjection,
5722        };
5723        let head = match weights.get(&head_id) {
5724            Some(tensor) => tensor_checked(&head_id, tensor, &[vocab, hidden])?,
5725            None => model_output,
5726        };
5727        let mut logits = linear(&final_hidden, head, tokens, hidden, vocab);
5728        apply_logits_transforms(&mut logits, vocab, &plan.logits);
5729        source_hidden = hidden_next.clone();
5730        outputs.push(ReferenceMtpOutput {
5731            depth: block.depth,
5732            logits,
5733            hidden: hidden_next,
5734            state,
5735        });
5736    }
5737    Ok(outputs)
5738}
5739
5740fn mla_attention(
5741    layer: u32,
5742    plan: &memra_gguf::model_plan::MlaAttentionPlan,
5743    epsilon: f32,
5744    weights: &ReferenceWeights,
5745    x: &[f32],
5746    tokens: usize,
5747    hidden: usize,
5748) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
5749    if let memra_gguf::model_plan::MlaAttentionPlan::CompressedKv { .. } = plan {
5750        return compressed_mla_attention(layer, plan, epsilon, weights, x, tokens, hidden);
5751    }
5752    let memra_gguf::model_plan::MlaAttentionPlan::LatentKv {
5753        query_heads,
5754        q_lora_rank,
5755        kv_lora_rank,
5756        qk_head_dim,
5757        rope_head_dim,
5758        value_head_dim,
5759        rope,
5760        sparse_index,
5761    } = plan.clone()
5762    else {
5763        return Err(ReferenceError::UnsupportedOperation {
5764            layer: Some(layer),
5765            operation: "compressed-KV MLA",
5766        });
5767    };
5768    // Per-token indexers execute only through full-selection equivalence; the k-pool
5769    // indexer (glm5_next) selects for real and is scored after q_resid exists.
5770    let plain_sparse_top_k = match &sparse_index {
5771        memra_gguf::model_plan::SparseIndexPlan::None
5772        | memra_gguf::model_plan::SparseIndexPlan::Own { kpool: Some(_), .. } => None,
5773        memra_gguf::model_plan::SparseIndexPlan::Own {
5774            top_k, kpool: None, ..
5775        }
5776        | memra_gguf::model_plan::SparseIndexPlan::SharedFromPrevious { top_k } => {
5777            Some(*top_k as usize)
5778        }
5779    };
5780    if plain_sparse_top_k.is_some_and(|top_k| tokens > top_k) {
5781        return Err(ReferenceError::UnsupportedOperation {
5782            layer: Some(layer),
5783            operation: "sparse MLA selection beyond full-selection equivalence",
5784        });
5785    }
5786    let heads = query_heads as usize;
5787    let q_rank = q_lora_rank as usize;
5788    let kv_rank = kv_lora_rank as usize;
5789    let qk_dim = qk_head_dim as usize;
5790    let rope_dim = rope_head_dim as usize;
5791    let nope_dim = qk_dim - rope_dim;
5792    let value_dim = value_head_dim as usize;
5793    let latent_dim = kv_rank + rope_dim;
5794
5795    let q_down = linear(
5796        x,
5797        tensor(
5798            weights,
5799            &layer_id(layer, LayerTensor::MlaQueryDown),
5800            &[q_rank, hidden],
5801        )?,
5802        tokens,
5803        hidden,
5804        q_rank,
5805    );
5806    let q_down = rms_norm(
5807        &q_down,
5808        tokens,
5809        q_rank,
5810        tensor(
5811            weights,
5812            &layer_id(layer, LayerTensor::MlaQueryDownNorm),
5813            &[q_rank],
5814        )?,
5815        epsilon,
5816    );
5817    // q_down is q_resid = q_a_layernorm(q_a_proj(x)): it feeds both the MLA query
5818    // up-projection and the k-pool indexer.
5819    let allowed_mask = match &sparse_index {
5820        memra_gguf::model_plan::SparseIndexPlan::Own {
5821            heads: index_heads,
5822            head_dim: index_dim,
5823            top_k,
5824            kpool: Some(kpool),
5825        } => {
5826            let allowed = kpool_allowed_tokens(
5827                layer,
5828                *index_heads as usize,
5829                *index_dim as usize,
5830                *top_k as usize,
5831                kpool,
5832                weights,
5833                x,
5834                &q_down,
5835                tokens,
5836                hidden,
5837                q_rank,
5838            )?;
5839            let mut mask = vec![false; tokens * tokens];
5840            for (token, sources) in allowed.iter().enumerate() {
5841                for &source in sources {
5842                    mask[token * tokens + source] = true;
5843                }
5844            }
5845            Some(mask)
5846        }
5847        _ => None,
5848    };
5849    let query = linear(
5850        &q_down,
5851        tensor(
5852            weights,
5853            &layer_id(layer, LayerTensor::MlaQueryUp),
5854            &[heads * qk_dim, q_rank],
5855        )?,
5856        tokens,
5857        q_rank,
5858        heads * qk_dim,
5859    );
5860    let latent_raw = linear(
5861        x,
5862        tensor(
5863            weights,
5864            &layer_id(layer, LayerTensor::MlaKvDown),
5865            &[latent_dim, hidden],
5866        )?,
5867        tokens,
5868        hidden,
5869        latent_dim,
5870    );
5871    let kv_norm = tensor(
5872        weights,
5873        &layer_id(layer, LayerTensor::MlaKvDownNorm),
5874        &[kv_rank],
5875    )?;
5876    let mut latent = latent_raw;
5877    for token in 0..tokens {
5878        let offset = token * latent_dim;
5879        let normalized = rms_norm(
5880            &latent[offset..offset + kv_rank],
5881            1,
5882            kv_rank,
5883            kv_norm,
5884            epsilon,
5885        );
5886        latent[offset..offset + kv_rank].copy_from_slice(&normalized);
5887    }
5888
5889    let mut query_nope = vec![0.0; tokens * heads * nope_dim];
5890    let mut query_rope = vec![0.0; tokens * heads * rope_dim];
5891    for token in 0..tokens {
5892        for head in 0..heads {
5893            let source = (token * heads + head) * qk_dim;
5894            let nope_target = (token * heads + head) * nope_dim;
5895            let rope_target = (token * heads + head) * rope_dim;
5896            query_nope[nope_target..nope_target + nope_dim]
5897                .copy_from_slice(&query[source..source + nope_dim]);
5898            query_rope[rope_target..rope_target + rope_dim]
5899                .copy_from_slice(&query[source + nope_dim..source + qk_dim]);
5900        }
5901    }
5902    let (rope_factors, rope_mscale) = rope_factor_values(&rope, weights)?;
5903    apply_rope(
5904        &mut query_rope,
5905        tokens,
5906        heads,
5907        rope_dim,
5908        rope.dimensions as usize,
5909        rope.base,
5910        rope_factors.as_deref(),
5911        rope_mscale,
5912    );
5913    let mut key_rope = vec![0.0; tokens * rope_dim];
5914    for token in 0..tokens {
5915        key_rope[token * rope_dim..(token + 1) * rope_dim]
5916            .copy_from_slice(&latent[token * latent_dim + kv_rank..(token + 1) * latent_dim]);
5917    }
5918    apply_rope(
5919        &mut key_rope,
5920        tokens,
5921        1,
5922        rope_dim,
5923        rope.dimensions as usize,
5924        rope.base,
5925        rope_factors.as_deref(),
5926        rope_mscale,
5927    );
5928    for token in 0..tokens {
5929        latent[token * latent_dim + kv_rank..(token + 1) * latent_dim]
5930            .copy_from_slice(&key_rope[token * rope_dim..(token + 1) * rope_dim]);
5931    }
5932
5933    // Contract layout: [head][kv_rank][nope] (see `deterministic_fixture`).
5934    let key_weight = tensor(
5935        weights,
5936        &layer_id(layer, LayerTensor::MlaKeyUp),
5937        &[heads, kv_rank, nope_dim],
5938    )?;
5939    let value_weight = tensor(
5940        weights,
5941        &layer_id(layer, LayerTensor::MlaValueUp),
5942        &[heads, value_dim, kv_rank],
5943    )?;
5944    let mut key_nope = vec![0.0; tokens * heads * nope_dim];
5945    let mut value = vec![0.0; tokens * heads * value_dim];
5946    for token in 0..tokens {
5947        let latent_row = &latent[token * latent_dim..token * latent_dim + kv_rank];
5948        for head in 0..heads {
5949            for out in 0..nope_dim {
5950                for rank in 0..kv_rank {
5951                    key_nope[(token * heads + head) * nope_dim + out] +=
5952                        latent_row[rank] * key_weight[(head * kv_rank + rank) * nope_dim + out];
5953                }
5954            }
5955            for out in 0..value_dim {
5956                for rank in 0..kv_rank {
5957                    value[(token * heads + head) * value_dim + out] +=
5958                        latent_row[rank] * value_weight[(head * value_dim + out) * kv_rank + rank];
5959                }
5960            }
5961        }
5962    }
5963    let mut attended = vec![0.0; tokens * heads * value_dim];
5964    let scale = 1.0 / (qk_dim as f32).sqrt();
5965    for token in 0..tokens {
5966        for head in 0..heads {
5967            let mut scores = Vec::with_capacity(token + 1);
5968            for source in 0..=token {
5969                // The indexer's allowed set masks keys exactly like the eager
5970                // additive -inf mask built from topk_indices.
5971                if allowed_mask
5972                    .as_ref()
5973                    .is_some_and(|mask| !mask[token * tokens + source])
5974                {
5975                    scores.push(f32::NEG_INFINITY);
5976                    continue;
5977                }
5978                let mut score = 0.0;
5979                for dim in 0..nope_dim {
5980                    score += query_nope[(token * heads + head) * nope_dim + dim]
5981                        * key_nope[(source * heads + head) * nope_dim + dim];
5982                }
5983                for dim in 0..rope_dim {
5984                    score += query_rope[(token * heads + head) * rope_dim + dim]
5985                        * key_rope[source * rope_dim + dim];
5986                }
5987                scores.push(score * scale);
5988            }
5989            softmax_in_place(&mut scores);
5990            for (source, probability) in scores.into_iter().enumerate() {
5991                for dim in 0..value_dim {
5992                    attended[(token * heads + head) * value_dim + dim] +=
5993                        probability * value[(source * heads + head) * value_dim + dim];
5994                }
5995            }
5996        }
5997    }
5998    let output = linear(
5999        &attended,
6000        tensor(
6001            weights,
6002            &layer_id(layer, LayerTensor::MlaOutput),
6003            &[hidden, heads * value_dim],
6004        )?,
6005        tokens,
6006        heads * value_dim,
6007        hidden,
6008    );
6009    Ok((
6010        output,
6011        ReferenceLayerState::LatentKv {
6012            rows: latent,
6013            tokens,
6014            width: latent_dim,
6015        },
6016    ))
6017}
6018
6019/// K-pool compressed indexer selection (Glm5NextTextIndexer.forward), single-sequence
6020/// causal case: every token is a valid key, so pooling starts at index 0 and only
6021/// causality masks candidates. Returns the allowed source-token set per query,
6022/// sorted ascending.
6023///
6024/// PUBLIC because it is the CUDA k-pool indexer's oracle: `memra-engine`'s
6025/// `tests/glm5_kpool_indexer_gpu.rs` compares the device's selected index sets against this
6026/// function on identical inputs. Its scope line above is part of the contract — a padded or
6027/// batched caller is outside it.
6028#[allow(clippy::too_many_arguments)]
6029pub fn kpool_allowed_tokens(
6030    layer: u32,
6031    index_heads: usize,
6032    index_dim: usize,
6033    top_k: usize,
6034    kpool: &memra_gguf::model_plan::KpoolPlan,
6035    weights: &ReferenceWeights,
6036    x: &[f32],
6037    q_resid: &[f32],
6038    tokens: usize,
6039    hidden: usize,
6040    q_rank: usize,
6041) -> Result<Vec<Vec<usize>>, ReferenceError> {
6042    let pool = kpool.pool as usize;
6043    if index_heads == 0 || index_dim == 0 || pool == 0 {
6044        return Err(ReferenceError::InvalidPlan {
6045            layer: Some(layer),
6046            reason: "k-pool sparse index requires positive heads, head_dim, and pool",
6047        });
6048    }
6049    let q = linear(
6050        q_resid,
6051        tensor(
6052            weights,
6053            &layer_id(layer, LayerTensor::SparseQuery),
6054            &[index_heads * index_dim, q_rank],
6055        )?,
6056        tokens,
6057        q_rank,
6058        index_heads * index_dim,
6059    );
6060    let key = layer_norm(
6061        &linear(
6062            x,
6063            tensor(
6064                weights,
6065                &layer_id(layer, LayerTensor::SparseKey),
6066                &[index_dim, hidden],
6067            )?,
6068            tokens,
6069            hidden,
6070            index_dim,
6071        ),
6072        tokens,
6073        index_dim,
6074        tensor(
6075            weights,
6076            &layer_id(layer, LayerTensor::SparseKeyNorm),
6077            &[index_dim],
6078        )?,
6079        tensor(
6080            weights,
6081            &layer_id(layer, LayerTensor::SparseKeyNormBias),
6082            &[index_dim],
6083        )?,
6084    );
6085    let gate_scores = linear(
6086        x,
6087        tensor(
6088            weights,
6089            &layer_id(layer, LayerTensor::SparseCompressorGate),
6090            &[index_dim, hidden],
6091        )?,
6092        tokens,
6093        hidden,
6094        index_dim,
6095    );
6096    let ape = tensor(
6097        weights,
6098        &layer_id(layer, LayerTensor::SparseCompressorPosition),
6099        &[pool, index_dim],
6100    )?;
6101    // Only COMPLETE pools are candidates; the incomplete tail never scores. Each
6102    // channel takes its own softmax over the pool members (gate score + APE).
6103    let pools = tokens / pool;
6104    let mut pool_keys = vec![0.0f32; pools * index_dim];
6105    for pool_index in 0..pools {
6106        for channel in 0..index_dim {
6107            let mut logits = Vec::with_capacity(pool);
6108            for slot in 0..pool {
6109                logits.push(
6110                    gate_scores[(pool_index * pool + slot) * index_dim + channel]
6111                        + ape[slot * index_dim + channel],
6112                );
6113            }
6114            softmax_in_place(&mut logits);
6115            let mut pooled = 0.0;
6116            for slot in 0..pool {
6117                pooled += logits[slot] * key[(pool_index * pool + slot) * index_dim + channel];
6118            }
6119            pool_keys[pool_index * index_dim + channel] = pooled;
6120        }
6121    }
6122    let mut head_weights = linear(
6123        x,
6124        tensor(
6125            weights,
6126            &layer_id(layer, LayerTensor::SparseProjection),
6127            &[index_heads, hidden],
6128        )?,
6129        tokens,
6130        hidden,
6131        index_heads,
6132    );
6133    let head_scale = (index_heads as f32).powf(-0.5);
6134    for value in &mut head_weights {
6135        *value *= head_scale;
6136    }
6137    // Same scale convention as the per-token DSA indexer: relu(q . k * hd^-0.5).
6138    let softmax_scale = (index_dim as f32).powf(-0.5);
6139    let select_k = (top_k / pool).min(pools);
6140    let mut allowed = Vec::with_capacity(tokens);
6141    for token in 0..tokens {
6142        // A pool is selectable only when its final token index is <= the query.
6143        let visible_pools = ((token + 1) / pool).min(pools);
6144        let mut scored: Vec<(usize, f32)> = (0..visible_pools)
6145            .map(|pool_index| {
6146                let mut score = 0.0f32;
6147                for head in 0..index_heads {
6148                    let mut dot = 0.0f32;
6149                    for dim in 0..index_dim {
6150                        dot += q[(token * index_heads + head) * index_dim + dim]
6151                            * pool_keys[pool_index * index_dim + dim];
6152                    }
6153                    score +=
6154                        (dot * softmax_scale).max(0.0) * head_weights[token * index_heads + head];
6155                }
6156                (pool_index, score)
6157            })
6158            .collect();
6159        scored.sort_by(|left, right| {
6160            right
6161                .1
6162                .partial_cmp(&left.1)
6163                .unwrap_or(std::cmp::Ordering::Equal)
6164                .then(left.0.cmp(&right.0))
6165        });
6166        let mut selected: Vec<usize> = Vec::new();
6167        for &(pool_index, _) in scored.iter().take(select_k) {
6168            selected.extend(pool_index * pool..(pool_index + 1) * pool);
6169        }
6170        if kpool.always_select_tail {
6171            // The current incomplete tail: the visible tokens past the last complete
6172            // visible pool (at most pool - 1 of them), always <= the query index.
6173            let visible = token + 1;
6174            let tail = visible % pool;
6175            selected.extend(visible - tail..visible);
6176        }
6177        if selected.is_empty() {
6178            // always_select_tail=false leaves early queries (before the first
6179            // complete pool) with no candidates; the reference would emit NaN rows.
6180            return Err(ReferenceError::InvalidPlan {
6181                layer: Some(layer),
6182                reason: "k-pool selection produced an empty candidate set for a query",
6183            });
6184        }
6185        selected.sort_unstable();
6186        allowed.push(selected);
6187    }
6188    Ok(allowed)
6189}
6190
6191#[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
6192fn compressed_mla_attention(
6193    layer: u32,
6194    plan: &memra_gguf::model_plan::MlaAttentionPlan,
6195    epsilon: f32,
6196    weights: &ReferenceWeights,
6197    x: &[f32],
6198    tokens: usize,
6199    hidden: usize,
6200) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
6201    use memra_gguf::dsv4_forward::{
6202        ActQuantVariant, IndexerW, apply_rope as apply_dsv4_rope, matmul, precompute_freqs_cis,
6203        rmsnorm,
6204    };
6205    use memra_gguf::model_plan::{MlaAttentionPlan, RopeFactors, SparseIndexPlan};
6206
6207    let MlaAttentionPlan::CompressedKv {
6208        query_heads,
6209        q_lora_rank,
6210        latent_head_dim,
6211        rope_head_dim,
6212        output_lora_rank,
6213        output_groups,
6214        window,
6215        rope,
6216        compressor,
6217        sparse_index,
6218    } = plan
6219    else {
6220        unreachable!()
6221    };
6222    let heads = *query_heads as usize;
6223    let q_rank = *q_lora_rank as usize;
6224    let head_dim = *latent_head_dim as usize;
6225    let rope_dim = *rope_head_dim as usize;
6226    let output_rank = *output_lora_rank as usize;
6227    let groups = *output_groups as usize;
6228    let window = *window as usize;
6229    if heads == 0
6230        || q_rank == 0
6231        || head_dim == 0
6232        || rope_dim == 0
6233        || rope_dim > head_dim
6234        || !(head_dim - rope_dim).is_multiple_of(64)
6235        || groups == 0
6236        || heads % groups != 0
6237        || window == 0
6238    {
6239        return Err(ReferenceError::InvalidPlan {
6240            layer: Some(layer),
6241            reason: "compressed attention has invalid reference geometry",
6242        });
6243    }
6244    let (original_context, factor, beta_fast, beta_slow) = match rope.factors {
6245        RopeFactors::None => (0, 1.0, 32.0, 1.0),
6246        RopeFactors::Yarn {
6247            factor,
6248            original_context,
6249            beta_fast,
6250            beta_slow,
6251        } => (original_context, factor, beta_fast, beta_slow),
6252        _ => {
6253            return Err(ReferenceError::InvalidPlan {
6254                layer: Some(layer),
6255                reason: "compressed attention requires plain or YaRN RoPE",
6256            });
6257        }
6258    };
6259    let frequencies = precompute_freqs_cis(
6260        rope_dim,
6261        tokens.max(1),
6262        original_context,
6263        rope.base,
6264        factor,
6265        beta_fast,
6266        beta_slow,
6267    );
6268    let positions: Vec<usize> = (0..tokens).collect();
6269
6270    let query_low_rank = rmsnorm(
6271        &matmul(
6272            x,
6273            tokens,
6274            hidden,
6275            tensor(
6276                weights,
6277                &layer_id(layer, LayerTensor::MlaQueryDown),
6278                &[q_rank, hidden],
6279            )?,
6280            q_rank,
6281        ),
6282        tensor(
6283            weights,
6284            &layer_id(layer, LayerTensor::MlaQueryDownNorm),
6285            &[q_rank],
6286        )?,
6287        epsilon,
6288    );
6289    let mut query = matmul(
6290        &query_low_rank,
6291        tokens,
6292        q_rank,
6293        tensor(
6294            weights,
6295            &layer_id(layer, LayerTensor::MlaQueryUp),
6296            &[heads * head_dim, q_rank],
6297        )?,
6298        heads * head_dim,
6299    );
6300    for head in query.chunks_exact_mut(head_dim) {
6301        let mean_square = head
6302            .iter()
6303            .map(|value| (*value as f64) * (*value as f64))
6304            .sum::<f64>()
6305            / head_dim as f64;
6306        let scale = 1.0 / (mean_square as f32 + epsilon).sqrt();
6307        for value in head {
6308            *value *= scale;
6309        }
6310    }
6311    apply_dsv4_rope(
6312        &mut query,
6313        tokens,
6314        heads,
6315        head_dim,
6316        rope_dim,
6317        &frequencies,
6318        &positions,
6319        false,
6320    );
6321
6322    let mut key_value = rmsnorm(
6323        &matmul(
6324            x,
6325            tokens,
6326            hidden,
6327            tensor(
6328                weights,
6329                &layer_id(layer, LayerTensor::MlaKvDown),
6330                &[head_dim, hidden],
6331            )?,
6332            head_dim,
6333        ),
6334        tensor(
6335            weights,
6336            &layer_id(layer, LayerTensor::MlaKvDownNorm),
6337            &[head_dim],
6338        )?,
6339        epsilon,
6340    );
6341    apply_dsv4_rope(
6342        &mut key_value,
6343        tokens,
6344        1,
6345        head_dim,
6346        rope_dim,
6347        &frequencies,
6348        &positions,
6349        false,
6350    );
6351    for row in key_value.chunks_exact_mut(head_dim) {
6352        memra_gguf::dsv4_forward::act_quant(
6353            &mut row[..head_dim - rope_dim],
6354            64,
6355            ActQuantVariant::RefFp8Round,
6356        );
6357    }
6358
6359    let (mut indices, mut slots) = memra_gguf::dsv4_forward::window_topk_idxs(window, tokens);
6360    let mut key_value_rows = tokens;
6361    let mut compressed_tokens = 0;
6362    if let Some(compressor_plan) = compressor {
6363        let ratio = compressor_plan.ratio as usize;
6364        let compressor = reference_compressor(
6365            weights,
6366            layer,
6367            hidden,
6368            head_dim,
6369            ratio,
6370            compressor_plan.latent_dim as usize,
6371            false,
6372        )?;
6373        let (compressed_indices, compressed_slots) = match sparse_index {
6374            SparseIndexPlan::None => {
6375                memra_gguf::dsv4_forward::compress_topk_idxs(ratio, tokens, tokens)
6376            }
6377            SparseIndexPlan::Own {
6378                heads: index_heads,
6379                head_dim: index_dim,
6380                top_k,
6381                kpool,
6382            } => {
6383                // K-pool scoring is a LatentKv (glm5_next) program; dsv4 compiles None.
6384                if kpool.is_some() {
6385                    return Err(ReferenceError::UnsupportedOperation {
6386                        layer: Some(layer),
6387                        operation: "k-pool sparse index on compressed attention",
6388                    });
6389                }
6390                let index_heads = *index_heads as usize;
6391                let index_dim = *index_dim as usize;
6392                if index_dim < rope_dim
6393                    || !index_dim.is_multiple_of(32)
6394                    || !index_dim.is_power_of_two()
6395                {
6396                    return Err(ReferenceError::InvalidPlan {
6397                        layer: Some(layer),
6398                        reason: "compressed sparse index has invalid head geometry",
6399                    });
6400                }
6401                let indexer = IndexerW {
6402                    wq_b: tensor(
6403                        weights,
6404                        &layer_id(layer, LayerTensor::SparseQuery),
6405                        &[index_heads * index_dim, q_rank],
6406                    )?
6407                    .to_vec(),
6408                    weights_proj: tensor(
6409                        weights,
6410                        &layer_id(layer, LayerTensor::SparseProjection),
6411                        &[index_heads, hidden],
6412                    )?
6413                    .to_vec(),
6414                    compressor: reference_compressor(
6415                        weights,
6416                        layer,
6417                        hidden,
6418                        index_dim,
6419                        ratio,
6420                        2 * index_dim,
6421                        true,
6422                    )?,
6423                    heads: index_heads,
6424                    hd: index_dim,
6425                    topk: *top_k as usize,
6426                };
6427                let output = indexer.forward(
6428                    x,
6429                    &query_low_rank,
6430                    tokens,
6431                    hidden,
6432                    q_rank,
6433                    tokens,
6434                    &frequencies,
6435                    rope_dim,
6436                    epsilon,
6437                    ActQuantVariant::RefFp8Round,
6438                    false,
6439                );
6440                (output.idxs, output.slots)
6441            }
6442            SparseIndexPlan::SharedFromPrevious { .. } => {
6443                return Err(ReferenceError::UnsupportedOperation {
6444                    layer: Some(layer),
6445                    operation: "shared compressed sparse-index execution",
6446                });
6447            }
6448        };
6449        if compressed_slots > 0 {
6450            let mut merged = vec![-1; tokens * (slots + compressed_slots)];
6451            for token in 0..tokens {
6452                merged[token * (slots + compressed_slots)
6453                    ..token * (slots + compressed_slots) + slots]
6454                    .copy_from_slice(&indices[token * slots..(token + 1) * slots]);
6455                merged[token * (slots + compressed_slots) + slots
6456                    ..(token + 1) * (slots + compressed_slots)]
6457                    .copy_from_slice(
6458                        &compressed_indices
6459                            [token * compressed_slots..(token + 1) * compressed_slots],
6460                    );
6461            }
6462            indices = merged;
6463            slots += compressed_slots;
6464        }
6465        if let Some((compressed, count)) = compressor.forward(
6466            x,
6467            tokens,
6468            hidden,
6469            &frequencies,
6470            rope_dim,
6471            epsilon,
6472            ActQuantVariant::RefFp8Round,
6473        ) {
6474            key_value.extend_from_slice(&compressed);
6475            key_value_rows += count;
6476            compressed_tokens = count;
6477        }
6478    }
6479
6480    let sink = tensor(
6481        weights,
6482        &layer_id(layer, LayerTensor::AttentionSink),
6483        &[heads],
6484    )?;
6485    let attention_scale = (head_dim as f64).powf(-0.5) as f32;
6486    let mut attended = vec![0.0; tokens * heads * head_dim];
6487    for token in 0..tokens {
6488        let selected = &indices[token * slots..(token + 1) * slots];
6489        memra_gguf::dsv4_decode::sparse_attn_query(
6490            &query[token * heads * head_dim..(token + 1) * heads * head_dim],
6491            heads,
6492            head_dim,
6493            selected,
6494            |index| &key_value[index * head_dim..(index + 1) * head_dim],
6495            sink,
6496            attention_scale,
6497            &mut attended[token * heads * head_dim..(token + 1) * heads * head_dim],
6498        );
6499    }
6500    apply_dsv4_rope(
6501        &mut attended,
6502        tokens,
6503        heads,
6504        head_dim,
6505        rope_dim,
6506        &frequencies,
6507        &positions,
6508        true,
6509    );
6510
6511    let group_width = heads / groups * head_dim;
6512    let output_down = tensor(
6513        weights,
6514        &layer_id(layer, LayerTensor::MlaOutputDown),
6515        &[groups * output_rank, group_width],
6516    )?;
6517    let mut grouped = vec![0.0; tokens * groups * output_rank];
6518    for token in 0..tokens {
6519        for group in 0..groups {
6520            let source = &attended[token * heads * head_dim + group * group_width
6521                ..token * heads * head_dim + (group + 1) * group_width];
6522            let group_weight = &output_down
6523                [group * output_rank * group_width..(group + 1) * output_rank * group_width];
6524            for rank in 0..output_rank {
6525                grouped[(token * groups + group) * output_rank + rank] =
6526                    memra_gguf::dsv4_forward::dot(
6527                        source,
6528                        &group_weight[rank * group_width..(rank + 1) * group_width],
6529                    );
6530            }
6531        }
6532    }
6533    let output = matmul(
6534        &grouped,
6535        tokens,
6536        groups * output_rank,
6537        tensor(
6538            weights,
6539            &layer_id(layer, LayerTensor::MlaOutput),
6540            &[hidden, groups * output_rank],
6541        )?,
6542        hidden,
6543    );
6544    Ok((
6545        output,
6546        ReferenceLayerState::CompressedAttention {
6547            rows: key_value,
6548            tokens: key_value_rows,
6549            width: head_dim,
6550            window,
6551            compressed_tokens,
6552        },
6553    ))
6554}
6555
6556#[allow(clippy::too_many_arguments)]
6557fn reference_compressor(
6558    weights: &ReferenceWeights,
6559    layer: u32,
6560    hidden: usize,
6561    output_dim: usize,
6562    ratio: usize,
6563    latent: usize,
6564    sparse: bool,
6565) -> Result<memra_gguf::dsv4_forward::CompressorW, ReferenceError> {
6566    let (key_value, gate, norm, position) = if sparse {
6567        (
6568            LayerTensor::SparseCompressorKeyValue,
6569            LayerTensor::SparseCompressorGate,
6570            LayerTensor::SparseCompressorNorm,
6571            LayerTensor::SparseCompressorPosition,
6572        )
6573    } else {
6574        (
6575            LayerTensor::KvCompressorKeyValue,
6576            LayerTensor::KvCompressorGate,
6577            LayerTensor::KvCompressorNorm,
6578            LayerTensor::KvCompressorPosition,
6579        )
6580    };
6581    Ok(memra_gguf::dsv4_forward::CompressorW {
6582        ratio,
6583        d: output_dim,
6584        latent,
6585        overlap: ratio == 4,
6586        rotate: sparse,
6587        wkv: tensor(weights, &layer_id(layer, key_value), &[latent, hidden])?.to_vec(),
6588        wgate: tensor(weights, &layer_id(layer, gate), &[latent, hidden])?.to_vec(),
6589        norm_w: tensor(weights, &layer_id(layer, norm), &[output_dim])?.to_vec(),
6590        ape: tensor(weights, &layer_id(layer, position), &[ratio, latent])?.to_vec(),
6591    })
6592}
6593
6594fn gated_delta_net(
6595    layer: u32,
6596    plan: &memra_gguf::model_plan::GatedDeltaNetPlan,
6597    epsilon: f32,
6598    weights: &ReferenceWeights,
6599    x: &[f32],
6600    tokens: usize,
6601    hidden: usize,
6602) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
6603    let key_heads = plan.key_heads as usize;
6604    let value_heads = plan.value_heads as usize;
6605    let key_dim = plan.key_head_dim as usize;
6606    let value_dim = plan.value_head_dim as usize;
6607    let kernel = plan.conv_kernel as usize;
6608    if key_heads == 0 || value_heads == 0 || key_dim == 0 || value_dim == 0 || kernel == 0 {
6609        return Err(ReferenceError::InvalidPlan {
6610            layer: Some(layer),
6611            reason: "GDN dimensions must be positive",
6612        });
6613    }
6614    let key_width = key_heads * key_dim;
6615    let value_width = value_heads * value_dim;
6616    let conv_width = 2 * key_width + value_width;
6617    let qkv = linear(
6618        x,
6619        tensor(
6620            weights,
6621            &layer_id(layer, LayerTensor::GdnQkv),
6622            &[conv_width, hidden],
6623        )?,
6624        tokens,
6625        hidden,
6626        conv_width,
6627    );
6628    let gate = linear(
6629        x,
6630        tensor(
6631            weights,
6632            &layer_id(layer, LayerTensor::GdnGate),
6633            &[value_width, hidden],
6634        )?,
6635        tokens,
6636        hidden,
6637        value_width,
6638    );
6639    let beta_raw = linear(
6640        x,
6641        tensor(
6642            weights,
6643            &layer_id(layer, LayerTensor::GdnBeta),
6644            &[value_heads, hidden],
6645        )?,
6646        tokens,
6647        hidden,
6648        value_heads,
6649    );
6650    let alpha = linear(
6651        x,
6652        tensor(
6653            weights,
6654            &layer_id(layer, LayerTensor::GdnAlpha),
6655            &[value_heads, hidden],
6656        )?,
6657        tokens,
6658        hidden,
6659        value_heads,
6660    );
6661    let conv_weight = tensor(
6662        weights,
6663        &layer_id(layer, LayerTensor::GdnConv1d),
6664        &[conv_width, kernel],
6665    )?;
6666    let mut conv = vec![0.0; tokens * conv_width];
6667    let pad = kernel - 1;
6668    for token in 0..tokens {
6669        for channel in 0..conv_width {
6670            let mut sum = 0.0;
6671            for tap in 0..kernel {
6672                let source = token as isize - pad as isize + tap as isize;
6673                if source >= 0 {
6674                    sum += qkv[source as usize * conv_width + channel]
6675                        * conv_weight[channel * kernel + tap];
6676                }
6677            }
6678            conv[token * conv_width + channel] = silu(sum);
6679        }
6680    }
6681
6682    let mut query = vec![0.0; tokens * value_heads * key_dim];
6683    let mut key = vec![0.0; tokens * value_heads * key_dim];
6684    let mut value = vec![0.0; tokens * value_width];
6685    for token in 0..tokens {
6686        for value_head in 0..value_heads {
6687            let key_head = value_head % key_heads;
6688            let q_source = token * conv_width + key_head * key_dim;
6689            let k_source = token * conv_width + key_width + key_head * key_dim;
6690            let v_source = token * conv_width + 2 * key_width + value_head * value_dim;
6691            let q_target = (token * value_heads + value_head) * key_dim;
6692            let v_target = (token * value_heads + value_head) * value_dim;
6693            query[q_target..q_target + key_dim]
6694                .copy_from_slice(&conv[q_source..q_source + key_dim]);
6695            key[q_target..q_target + key_dim].copy_from_slice(&conv[k_source..k_source + key_dim]);
6696            value[v_target..v_target + value_dim]
6697                .copy_from_slice(&conv[v_source..v_source + value_dim]);
6698        }
6699    }
6700    l2_normalize_rows(&mut query, tokens * value_heads, key_dim, epsilon);
6701    l2_normalize_rows(&mut key, tokens * value_heads, key_dim, epsilon);
6702
6703    let a = tensor(weights, &layer_id(layer, LayerTensor::GdnA), &[value_heads])?;
6704    let dt = tensor(
6705        weights,
6706        &layer_id(layer, LayerTensor::GdnDtBias),
6707        &[value_heads],
6708    )?;
6709    let mut matrix = vec![0.0; value_heads * value_dim * key_dim];
6710    let mut mixed = vec![0.0; tokens * value_width];
6711    let scale = 1.0 / (key_dim as f32).sqrt();
6712    for token in 0..tokens {
6713        for head in 0..value_heads {
6714            let beta = sigmoid(beta_raw[token * value_heads + head]);
6715            let decay = (a[head] * softplus(alpha[token * value_heads + head] + dt[head])).exp();
6716            let q_offset = (token * value_heads + head) * key_dim;
6717            let v_offset = (token * value_heads + head) * value_dim;
6718            let state_offset = head * value_dim * key_dim;
6719            let mut next = matrix[state_offset..state_offset + value_dim * key_dim].to_vec();
6720            for value_index in 0..value_dim {
6721                let row = state_offset + value_index * key_dim;
6722                let mut state_key = 0.0;
6723                for key_index in 0..key_dim {
6724                    state_key += matrix[row + key_index] * key[q_offset + key_index];
6725                }
6726                let delta = (value[v_offset + value_index] - decay * state_key) * beta;
6727                let mut attended = 0.0;
6728                for key_index in 0..key_dim {
6729                    let updated =
6730                        decay * matrix[row + key_index] + key[q_offset + key_index] * delta;
6731                    next[value_index * key_dim + key_index] = updated;
6732                    attended += updated * query[q_offset + key_index];
6733                }
6734                mixed[v_offset + value_index] = attended * scale;
6735            }
6736            matrix[state_offset..state_offset + value_dim * key_dim].copy_from_slice(&next);
6737        }
6738    }
6739
6740    let norm = tensor(
6741        weights,
6742        &layer_id(layer, LayerTensor::GdnNorm),
6743        &[value_dim],
6744    )?;
6745    let normalized = rms_norm(&mixed, tokens * value_heads, value_dim, norm, epsilon);
6746    let mut gated = normalized;
6747    for index in 0..gated.len() {
6748        // qwen4_exp declares sigmoid here (config output_gate_type) — the ONE numeric
6749        // divergence from the qwen3_5 GDN program (SEMANTICS.md §GDN); every other
6750        // family is the silu arm.
6751        gated[index] *= match plan.gate_activation {
6752            GdnGateActivation::Silu => silu(gate[index]),
6753            GdnGateActivation::Sigmoid => sigmoid(gate[index]),
6754        };
6755    }
6756    let output = linear(
6757        &gated,
6758        tensor(
6759            weights,
6760            &layer_id(layer, LayerTensor::GdnOutput),
6761            &[hidden, value_width],
6762        )?,
6763        tokens,
6764        value_width,
6765        hidden,
6766    );
6767    let mut conv_state = vec![0.0; conv_width * pad];
6768    for channel in 0..conv_width {
6769        for index in 0..pad {
6770            let source = tokens as isize - pad as isize + index as isize;
6771            if source >= 0 {
6772                conv_state[channel * pad + index] = qkv[source as usize * conv_width + channel];
6773            }
6774        }
6775    }
6776    Ok((
6777        output,
6778        ReferenceLayerState::Recurrent {
6779            conv: conv_state,
6780            matrix,
6781            value_heads,
6782            key_head_dim: key_dim,
6783            value_head_dim: value_dim,
6784            conv_width,
6785        },
6786    ))
6787}
6788
6789#[allow(clippy::too_many_arguments)]
6790/// Kimi Delta Attention (recurrent_kimi_delta_attention + Glm5NextTextLinearAttention),
6791/// all f32, sequential over tokens. Only the lower-bound forget-gate branch exists:
6792/// GLM-5.3-Flash always configures `gate_lower_bound`, so the softplus branch of
6793/// Glm5NextTextForgetGate is dead for this model and deliberately not implemented.
6794/// GPU-parity seam: run ONE KDA layer's mixer over `x` (`[tokens, hidden]`, already
6795/// pre-attention-normed) and return its output plus the recurrent state it leaves behind.
6796///
6797/// This is the very `kimi_delta_net` the trunk executor dispatches — exposed so
6798/// `crates/memra-engine/tests/kda_fixture_gpu.rs` can gate the CUDA mixer against the pinned
6799/// reference without standing up a whole model (glm5_next's residual topology and MLA layers
6800/// are a different surface, and a mixer gate must not depend on them).
6801pub fn kimi_delta_net_layer(
6802    layer: u32,
6803    plan: &memra_gguf::model_plan::KimiDeltaNetPlan,
6804    epsilon: f32,
6805    weights: &ReferenceWeights,
6806    x: &[f32],
6807    tokens: usize,
6808    hidden: usize,
6809) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
6810    kimi_delta_net(layer, plan, epsilon, weights, x, tokens, hidden)
6811}
6812
6813fn kimi_delta_net(
6814    layer: u32,
6815    plan: &memra_gguf::model_plan::KimiDeltaNetPlan,
6816    epsilon: f32,
6817    weights: &ReferenceWeights,
6818    x: &[f32],
6819    tokens: usize,
6820    hidden: usize,
6821) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
6822    let heads = plan.num_heads as usize;
6823    let head_dim = plan.head_dim as usize;
6824    let kernel = plan.conv_kernel as usize;
6825    if heads == 0 || head_dim == 0 || kernel == 0 {
6826        return Err(ReferenceError::InvalidPlan {
6827            layer: Some(layer),
6828            reason: "KDA dimensions must be positive",
6829        });
6830    }
6831    let qkv = heads * head_dim;
6832    let conv_width = 3 * qkv;
6833    let project_and_convolve = |projection: LayerTensor,
6834                                conv: LayerTensor|
6835     -> Result<(Vec<f32>, Vec<f32>), ReferenceError> {
6836        let projected = linear(
6837            x,
6838            tensor(weights, &layer_id(layer, projection), &[qkv, hidden])?,
6839            tokens,
6840            hidden,
6841            qkv,
6842        );
6843        // The checkpoint splits the grouped causal conv into three per-plane
6844        // convs; applying each to its own plane is the fused conv exactly.
6845        let conv_weight = tensor(weights, &layer_id(layer, conv), &[qkv, kernel])?;
6846        let mut convolved = vec![0.0; tokens * qkv];
6847        for token in 0..tokens {
6848            for channel in 0..qkv {
6849                let mut sum = 0.0;
6850                for tap in 0..kernel {
6851                    let source = token as isize - (kernel - 1) as isize + tap as isize;
6852                    if source >= 0 {
6853                        sum += projected[source as usize * qkv + channel]
6854                            * conv_weight[channel * kernel + tap];
6855                    }
6856                }
6857                convolved[token * qkv + channel] = silu(sum);
6858            }
6859        }
6860        Ok((projected, convolved))
6861    };
6862    let (q_raw, mut query) =
6863        project_and_convolve(LayerTensor::KdaQuery, LayerTensor::KdaQueryConv)?;
6864    let (k_raw, mut key) = project_and_convolve(LayerTensor::KdaKey, LayerTensor::KdaKeyConv)?;
6865    let (v_raw, value) = project_and_convolve(LayerTensor::KdaValue, LayerTensor::KdaValueConv)?;
6866    // FLA l2norm: x / sqrt(sum(x^2) + 1e-6) — the epsilon sits INSIDE the sqrt and is
6867    // fixed at 1e-6, independent of the layer epsilon.
6868    l2_normalize_rows(&mut query, tokens * heads, head_dim, 1e-6);
6869    l2_normalize_rows(&mut key, tokens * heads, head_dim, 1e-6);
6870    // Query scale head_dim^-0.5 applies AFTER the l2norm.
6871    let query_scale = 1.0 / (head_dim as f32).sqrt();
6872    for entry in &mut query {
6873        *entry *= query_scale;
6874    }
6875
6876    // Forget gate: g = gate_lower_bound * sigmoid(exp(A_log[head]) * (f_b(f_a(x)) + dt_bias)),
6877    // per channel (dt_bias has width qkv).
6878    let forget_down = linear(
6879        x,
6880        tensor(
6881            weights,
6882            &layer_id(layer, LayerTensor::KdaForgetDown),
6883            &[head_dim, hidden],
6884        )?,
6885        tokens,
6886        hidden,
6887        head_dim,
6888    );
6889    let mut forget = linear(
6890        &forget_down,
6891        tensor(
6892            weights,
6893            &layer_id(layer, LayerTensor::KdaForgetUp),
6894            &[qkv, head_dim],
6895        )?,
6896        tokens,
6897        head_dim,
6898        qkv,
6899    );
6900    let dt_bias = tensor(weights, &layer_id(layer, LayerTensor::KdaDtBias), &[qkv])?;
6901    let a_log = tensor(weights, &layer_id(layer, LayerTensor::KdaALog), &[heads])?;
6902    for token in 0..tokens {
6903        #[allow(clippy::needless_range_loop)]
6904        // allow: the explicit index loop keeps the offset arithmetic visible and aligned with the device-side indexing
6905        for head in 0..heads {
6906            let decay_rate = a_log[head].exp();
6907            for dim in 0..head_dim {
6908                let channel = head * head_dim + dim;
6909                let raw = forget[token * qkv + channel] + dt_bias[channel];
6910                forget[token * qkv + channel] = plan.gate_lower_bound * sigmoid(decay_rate * raw);
6911            }
6912        }
6913    }
6914    let beta_raw = linear(
6915        x,
6916        tensor(
6917            weights,
6918            &layer_id(layer, LayerTensor::KdaBeta),
6919            &[heads, hidden],
6920        )?,
6921        tokens,
6922        hidden,
6923        heads,
6924    );
6925
6926    // Recurrence (recurrent_kimi_delta_attention:477-489): state [heads, k_dim, v_dim];
6927    // exp(g) decays along the K dimension.
6928    let mut matrix = vec![0.0; heads * head_dim * head_dim];
6929    let mut core = vec![0.0; tokens * qkv];
6930    for token in 0..tokens {
6931        for head in 0..heads {
6932            let beta = sigmoid(beta_raw[token * heads + head]);
6933            let row_offset = (token * heads + head) * head_dim;
6934            let state_offset = head * head_dim * head_dim;
6935            for key_index in 0..head_dim {
6936                let decay = forget[token * qkv + head * head_dim + key_index].exp();
6937                let state_row = state_offset + key_index * head_dim;
6938                for value_index in 0..head_dim {
6939                    matrix[state_row + value_index] *= decay;
6940                }
6941            }
6942            let mut delta = vec![0.0; head_dim];
6943            for value_index in 0..head_dim {
6944                let mut memory = 0.0;
6945                for key_index in 0..head_dim {
6946                    memory += matrix[state_offset + key_index * head_dim + value_index]
6947                        * key[row_offset + key_index];
6948                }
6949                delta[value_index] = (value[row_offset + value_index] - memory) * beta;
6950            }
6951            for key_index in 0..head_dim {
6952                let state_row = state_offset + key_index * head_dim;
6953                for value_index in 0..head_dim {
6954                    matrix[state_row + value_index] +=
6955                        key[row_offset + key_index] * delta[value_index];
6956                }
6957            }
6958            for value_index in 0..head_dim {
6959                let mut attended = 0.0;
6960                for key_index in 0..head_dim {
6961                    attended += matrix[state_offset + key_index * head_dim + value_index]
6962                        * query[row_offset + key_index];
6963                }
6964                core[row_offset + value_index] = attended;
6965            }
6966        }
6967    }
6968
6969    // Output: sigmoid-gated fp32 RMSNorm over head_dim (o_norm uses the layer's
6970    // rms_norm_eps), gate = g_b(g_a(x)); then o_proj.
6971    let gate_down = linear(
6972        x,
6973        tensor(
6974            weights,
6975            &layer_id(layer, LayerTensor::KdaGateDown),
6976            &[head_dim, hidden],
6977        )?,
6978        tokens,
6979        hidden,
6980        head_dim,
6981    );
6982    let gate = linear(
6983        &gate_down,
6984        tensor(
6985            weights,
6986            &layer_id(layer, LayerTensor::KdaGateUp),
6987            &[qkv, head_dim],
6988        )?,
6989        tokens,
6990        head_dim,
6991        qkv,
6992    );
6993    let norm_weight = tensor(
6994        weights,
6995        &layer_id(layer, LayerTensor::KdaOutputNorm),
6996        &[head_dim],
6997    )?;
6998    let mut gated = rms_norm(&core, tokens * heads, head_dim, norm_weight, epsilon);
6999    for index in 0..gated.len() {
7000        gated[index] *= sigmoid(gate[index]);
7001    }
7002    let output = linear(
7003        &gated,
7004        tensor(
7005            weights,
7006            &layer_id(layer, LayerTensor::KdaOutput),
7007            &[hidden, qkv],
7008        )?,
7009        tokens,
7010        qkv,
7011        hidden,
7012    );
7013
7014    // Conv state stores the raw fused [q|k|v] pre-conv planes for the trailing
7015    // kernel-1 positions, mirroring the GDN layout.
7016    let pad = kernel - 1;
7017    let mut conv_state = vec![0.0; conv_width * pad];
7018    let planes = [&q_raw, &k_raw, &v_raw];
7019    for channel in 0..conv_width {
7020        let plane = channel / qkv;
7021        let plane_channel = channel % qkv;
7022        for index in 0..pad {
7023            let source = tokens as isize - pad as isize + index as isize;
7024            if source >= 0 {
7025                conv_state[channel * pad + index] =
7026                    planes[plane][source as usize * qkv + plane_channel];
7027            }
7028        }
7029    }
7030    Ok((
7031        output,
7032        ReferenceLayerState::Recurrent {
7033            conv: conv_state,
7034            matrix,
7035            value_heads: heads,
7036            key_head_dim: head_dim,
7037            value_head_dim: head_dim,
7038            conv_width,
7039        },
7040    ))
7041}
7042
7043#[allow(clippy::too_many_arguments)]
7044// allow: the parameter list mirrors the kernel/FFI/call contract; bundling into a struct is a refactor, not a lint fix
7045#[allow(clippy::manual_is_multiple_of)] // allow: divisor is runtime-derived; the modulo form keeps a zero divisor loud (a panic), where is_multiple_of would return false silently
7046fn full_attention(
7047    layer: u32,
7048    plan: &memra_gguf::model_plan::FullAttentionPlan,
7049    window: Option<usize>,
7050    norm_epsilon: f32,
7051    weights: &ReferenceWeights,
7052    x: &[f32],
7053    tokens: usize,
7054    hidden: usize,
7055    // qwen4_exp QSA: indexer visibility overlay, `[tokens, tokens]` row-major (query,
7056    // source); attention runs dense under causal AND selection (SEMANTICS.md §QSA).
7057    selection: Option<&[bool]>,
7058) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
7059    let query_heads = plan.query_heads as usize;
7060    let kv_heads = plan.kv_heads as usize;
7061    let key_dim = plan.key_head_dim as usize;
7062    let value_dim = plan.value_head_dim as usize;
7063    if query_heads == 0 || kv_heads == 0 || query_heads % kv_heads != 0 {
7064        return Err(ReferenceError::InvalidPlan {
7065            layer: Some(layer),
7066            reason: "query heads must be a positive multiple of KV heads",
7067        });
7068    }
7069    if selection.is_some_and(|selection| selection.len() != tokens * tokens) {
7070        return Err(ReferenceError::InvalidPlan {
7071            layer: Some(layer),
7072            reason: "attention selection mask does not match tokens x tokens",
7073        });
7074    }
7075    let fused = plan.output_gate == AttentionGateKind::FusedQ;
7076    let q_width = query_heads * key_dim;
7077    let q_projection_width = q_width * if fused { 2 } else { 1 };
7078    let k_width = kv_heads * key_dim;
7079    let v_width = kv_heads * value_dim;
7080    let q_weight = tensor(
7081        weights,
7082        &layer_id(layer, LayerTensor::Query),
7083        &[q_projection_width, hidden],
7084    )?;
7085    let k_weight = tensor(
7086        weights,
7087        &layer_id(layer, LayerTensor::Key),
7088        &[k_width, hidden],
7089    )?;
7090    let output_weight = tensor(
7091        weights,
7092        &layer_id(layer, LayerTensor::AttentionOutput),
7093        &[hidden, query_heads * value_dim],
7094    )?;
7095    let q_projected = linear(x, q_weight, tokens, hidden, q_projection_width);
7096    let mut query = vec![0.0; tokens * q_width];
7097    let mut fused_gate = None;
7098    if fused {
7099        let mut gate = vec![0.0; tokens * q_width];
7100        for token in 0..tokens {
7101            for head in 0..query_heads {
7102                let projected = token * q_projection_width + head * 2 * key_dim;
7103                let canonical = (token * query_heads + head) * key_dim;
7104                query[canonical..canonical + key_dim]
7105                    .copy_from_slice(&q_projected[projected..projected + key_dim]);
7106                gate[canonical..canonical + key_dim]
7107                    .copy_from_slice(&q_projected[projected + key_dim..projected + 2 * key_dim]);
7108            }
7109        }
7110        fused_gate = Some(gate);
7111    } else {
7112        query.copy_from_slice(&q_projected);
7113    }
7114    let mut key = linear(x, k_weight, tokens, hidden, k_width);
7115    let mut value = match plan.value_projection {
7116        ValueProjection::Separate => linear(
7117            x,
7118            tensor(
7119                weights,
7120                &layer_id(layer, LayerTensor::Value),
7121                &[v_width, hidden],
7122            )?,
7123            tokens,
7124            hidden,
7125            v_width,
7126        ),
7127        ValueProjection::ReuseKey => {
7128            if value_dim != key_dim {
7129                return Err(ReferenceError::InvalidPlan {
7130                    layer: Some(layer),
7131                    reason: "K-as-V requires equal key/value head widths",
7132                });
7133            }
7134            key.clone()
7135        }
7136    };
7137    apply_optional_head_norm(
7138        weights,
7139        layer_id(layer, LayerTensor::QueryNorm),
7140        &mut query,
7141        tokens * query_heads,
7142        key_dim,
7143        plan.qk_norm,
7144        norm_epsilon,
7145    )?;
7146    if plan.value_norm == ValueNorm::WeightlessRms {
7147        let ones = vec![1.0; value_dim];
7148        value = rms_norm(&value, tokens * kv_heads, value_dim, &ones, norm_epsilon);
7149    }
7150    apply_optional_head_norm(
7151        weights,
7152        layer_id(layer, LayerTensor::KeyNorm),
7153        &mut key,
7154        tokens * kv_heads,
7155        key_dim,
7156        plan.qk_norm,
7157        norm_epsilon,
7158    )?;
7159    let (rope_factors, rope_mscale) = rope_factor_values(&plan.rope, weights)?;
7160    apply_rope(
7161        &mut query,
7162        tokens,
7163        query_heads,
7164        key_dim,
7165        plan.rope.dimensions as usize,
7166        plan.rope.base,
7167        rope_factors.as_deref(),
7168        rope_mscale,
7169    );
7170    apply_rope(
7171        &mut key,
7172        tokens,
7173        kv_heads,
7174        key_dim,
7175        plan.rope.dimensions as usize,
7176        plan.rope.base,
7177        rope_factors.as_deref(),
7178        rope_mscale,
7179    );
7180
7181    let mut attended = vec![0.0; tokens * query_heads * value_dim];
7182    let scale = match plan.scale {
7183        AttentionScale::InverseSqrtKeyDim => 1.0 / (key_dim as f32).sqrt(),
7184        AttentionScale::Fixed(scale) => scale,
7185    };
7186    for token in 0..tokens {
7187        for head in 0..query_heads {
7188            let kv_head = head * kv_heads / query_heads;
7189            let first_source = window
7190                .map(|window| (token + 1).saturating_sub(window))
7191                .unwrap_or(0);
7192            let mut sources = Vec::with_capacity(token + 1 - first_source);
7193            let mut scores = Vec::with_capacity(token + 1 - first_source);
7194            for source in first_source..=token {
7195                if selection.is_some_and(|selection| !selection[token * tokens + source]) {
7196                    continue;
7197                }
7198                let mut score = 0.0;
7199                for dim in 0..key_dim {
7200                    score += query[(token * query_heads + head) * key_dim + dim]
7201                        * key[(source * kv_heads + kv_head) * key_dim + dim];
7202                }
7203                sources.push(source);
7204                scores.push(score * scale);
7205            }
7206            if scores.is_empty() {
7207                // The QSA tail rule guarantees every query keeps at least its own block's
7208                // incomplete tail; an empty row means a malformed selection mask.
7209                return Err(ReferenceError::InvalidPlan {
7210                    layer: Some(layer),
7211                    reason: "attention selection left a query with no visible source",
7212                });
7213            }
7214            softmax_in_place(&mut scores);
7215            for (index, probability) in scores.into_iter().enumerate() {
7216                let source = sources[index];
7217                for dim in 0..value_dim {
7218                    attended[(token * query_heads + head) * value_dim + dim] +=
7219                        probability * value[(source * kv_heads + kv_head) * value_dim + dim];
7220                }
7221            }
7222        }
7223    }
7224    if let Some(gate) = fused_gate {
7225        for token in 0..tokens {
7226            for head in 0..query_heads {
7227                for dim in 0..value_dim {
7228                    if dim >= key_dim {
7229                        return Err(ReferenceError::InvalidPlan {
7230                            layer: Some(layer),
7231                            reason: "fused attention gate requires value_dim <= key_dim",
7232                        });
7233                    }
7234                    attended[(token * query_heads + head) * value_dim + dim] *=
7235                        sigmoid(gate[(token * query_heads + head) * key_dim + dim]);
7236                }
7237            }
7238        }
7239    } else if plan.output_gate == AttentionGateKind::SeparateHead {
7240        let gate_weight = tensor(
7241            weights,
7242            &layer_id(layer, LayerTensor::AttentionGate),
7243            &[query_heads, hidden],
7244        )?;
7245        let gates = linear(x, gate_weight, tokens, hidden, query_heads);
7246        for token in 0..tokens {
7247            for head in 0..query_heads {
7248                let gate = sigmoid(gates[token * query_heads + head]);
7249                for dim in 0..value_dim {
7250                    attended[(token * query_heads + head) * value_dim + dim] *= gate;
7251                }
7252            }
7253        }
7254    }
7255    let state_start = window
7256        .map(|window| tokens.saturating_sub(window))
7257        .unwrap_or(0);
7258    let state_tokens = tokens - state_start;
7259    let state_key = key[state_start * k_width..].to_vec();
7260    let state_value = value[state_start * v_width..].to_vec();
7261    Ok((
7262        linear(
7263            &attended,
7264            output_weight,
7265            tokens,
7266            query_heads * value_dim,
7267            hidden,
7268        ),
7269        ReferenceLayerState::Kv {
7270            key: state_key,
7271            value: state_value,
7272            tokens: state_tokens,
7273            kv_heads,
7274            key_head_dim: key_dim,
7275            value_head_dim: value_dim,
7276            window,
7277        },
7278    ))
7279}
7280
7281fn dense_mlp(
7282    layer: u32,
7283    plan: &memra_gguf::model_plan::DenseMlpPlan,
7284    weights: &ReferenceWeights,
7285    x: &[f32],
7286    tokens: usize,
7287    hidden: usize,
7288) -> Result<Vec<f32>, ReferenceError> {
7289    let intermediate = plan.intermediate_size as usize;
7290    let gate = linear(
7291        x,
7292        tensor(
7293            weights,
7294            &layer_id(layer, LayerTensor::MlpGate),
7295            &[intermediate, hidden],
7296        )?,
7297        tokens,
7298        hidden,
7299        intermediate,
7300    );
7301    let up = linear(
7302        x,
7303        tensor(
7304            weights,
7305            &layer_id(layer, LayerTensor::MlpUp),
7306            &[intermediate, hidden],
7307        )?,
7308        tokens,
7309        hidden,
7310        intermediate,
7311    );
7312    let mut activated = vec![0.0; gate.len()];
7313    for index in 0..gate.len() {
7314        activated[index] = activate_pair(&plan.activation, gate[index], up[index], layer)?;
7315    }
7316    Ok(linear(
7317        &activated,
7318        tensor(
7319            weights,
7320            &layer_id(layer, LayerTensor::MlpDown),
7321            &[hidden, intermediate],
7322        )?,
7323        tokens,
7324        intermediate,
7325        hidden,
7326    ))
7327}
7328
7329#[allow(clippy::too_many_arguments)] // allow: the parameter list mirrors the kernel/FFI/call contract; bundling into a struct is a refactor, not a lint fix
7330fn moe_mlp(
7331    layer: u32,
7332    plan: &memra_gguf::model_plan::MoeMlpPlan,
7333    weights: &ReferenceWeights,
7334    x: &[f32],
7335    token_ids: &[u32],
7336    tokens: usize,
7337    hidden: usize,
7338    vocab: usize,
7339) -> Result<Vec<f32>, ReferenceError> {
7340    let experts = plan.expert_count as usize;
7341    let selected = plan.experts_per_token as usize;
7342    let intermediate = plan.expert_intermediate_size as usize;
7343    if selected == 0 || selected > experts {
7344        return Err(ReferenceError::InvalidPlan {
7345            layer: Some(layer),
7346            reason: "MoE top-k must be in 1..=expert_count",
7347        });
7348    }
7349    let router = tensor(
7350        weights,
7351        &layer_id(layer, LayerTensor::MoeRouter),
7352        &[experts, hidden],
7353    )?;
7354    let logits = linear(x, router, tokens, hidden, experts);
7355    let bias = if router_has_selection_bias(&plan.router) {
7356        Some(tensor(
7357            weights,
7358            &layer_id(layer, LayerTensor::MoeRouterBias),
7359            &[experts],
7360        )?)
7361    } else {
7362        None
7363    };
7364    let token_to_expert = if matches!(
7365        plan.router,
7366        memra_gguf::model_plan::RouterPlan::TokenIdHash { .. }
7367    ) {
7368        Some(tensor(
7369            weights,
7370            &layer_id(layer, LayerTensor::MoeTokenToExpert),
7371            &[vocab, selected],
7372        )?)
7373    } else {
7374        None
7375    };
7376    let gate_bank = tensor(
7377        weights,
7378        &layer_id(layer, LayerTensor::MoeExpertGateBank),
7379        &[experts, intermediate, hidden],
7380    )?;
7381    let up_bank = tensor(
7382        weights,
7383        &layer_id(layer, LayerTensor::MoeExpertUpBank),
7384        &[experts, intermediate, hidden],
7385    )?;
7386    let down_bank = tensor(
7387        weights,
7388        &layer_id(layer, LayerTensor::MoeExpertDownBank),
7389        &[experts, hidden, intermediate],
7390    )?;
7391    let mut output = vec![0.0; tokens * hidden];
7392    for token in 0..tokens {
7393        let forced_routes = token_to_expert
7394            .map(|table| {
7395                let token_id = token_ids[token] as usize;
7396                &table[token_id * selected..(token_id + 1) * selected]
7397            })
7398            .map(|row| {
7399                row.iter()
7400                    .map(|&value| {
7401                        if !value.is_finite()
7402                            || value < 0.0
7403                            || value.fract() != 0.0
7404                            || value as usize >= experts
7405                        {
7406                            return Err(ReferenceError::InvalidPlan {
7407                                layer: Some(layer),
7408                                reason: "token-id expert table contains an invalid expert id",
7409                            });
7410                        }
7411                        Ok(value as usize)
7412                    })
7413                    .collect::<Result<Vec<_>, _>>()
7414            })
7415            .transpose()?;
7416        let routes = route_experts(
7417            &plan.router,
7418            &logits[token * experts..(token + 1) * experts],
7419            bias,
7420            selected,
7421            forced_routes.as_deref(),
7422            layer,
7423        )?;
7424        if crate::hidden_trace::enabled() && token + 1 == tokens {
7425            crate::hidden_trace::emit_last_row(
7426                "router",
7427                layer as i64,
7428                1,
7429                experts,
7430                &logits[token * experts..(token + 1) * experts],
7431            );
7432            let mut route = Vec::with_capacity(routes.len() * 2);
7433            for (expert, weight) in &routes {
7434                route.push(*expert as f32);
7435                route.push(*weight);
7436            }
7437            crate::hidden_trace::emit_last_row("route", layer as i64, 1, route.len(), &route);
7438        }
7439        let input = &x[token * hidden..(token + 1) * hidden];
7440        for (expert, route_weight) in routes {
7441            let gate_offset = expert * intermediate * hidden;
7442            let down_offset = expert * hidden * intermediate;
7443            let mut activated = vec![0.0; intermediate];
7444            for row in 0..intermediate {
7445                let mut gate = 0.0;
7446                let mut up = 0.0;
7447                for column in 0..hidden {
7448                    gate += input[column] * gate_bank[gate_offset + row * hidden + column];
7449                    up += input[column] * up_bank[gate_offset + row * hidden + column];
7450                }
7451                activated[row] = activate_pair(&plan.activation, gate, up, layer)?;
7452            }
7453            for row in 0..hidden {
7454                let mut value = 0.0;
7455                for column in 0..intermediate {
7456                    value +=
7457                        activated[column] * down_bank[down_offset + row * intermediate + column];
7458                }
7459                output[token * hidden + row] += route_weight * value;
7460            }
7461        }
7462    }
7463
7464    if crate::hidden_trace::enabled() {
7465        crate::hidden_trace::emit_last_row("routed", layer as i64, tokens, hidden, &output);
7466    }
7467
7468    if let Some(shared) = plan.shared.as_ref() {
7469        let intermediate = shared.intermediate_size as usize;
7470        let gate = linear(
7471            x,
7472            tensor(
7473                weights,
7474                &layer_id(layer, LayerTensor::SharedMlpGate),
7475                &[intermediate, hidden],
7476            )?,
7477            tokens,
7478            hidden,
7479            intermediate,
7480        );
7481        let up = linear(
7482            x,
7483            tensor(
7484                weights,
7485                &layer_id(layer, LayerTensor::SharedMlpUp),
7486                &[intermediate, hidden],
7487            )?,
7488            tokens,
7489            hidden,
7490            intermediate,
7491        );
7492        let mut activated = vec![0.0; gate.len()];
7493        for index in 0..gate.len() {
7494            activated[index] = activate_pair(&plan.activation, gate[index], up[index], layer)?;
7495        }
7496        let mut shared_output = linear(
7497            &activated,
7498            tensor(
7499                weights,
7500                &layer_id(layer, LayerTensor::SharedMlpDown),
7501                &[hidden, intermediate],
7502            )?,
7503            tokens,
7504            intermediate,
7505            hidden,
7506        );
7507        if shared.gated {
7508            let gate_weight = tensor(
7509                weights,
7510                &layer_id(layer, LayerTensor::SharedMlpInputGate),
7511                &[hidden],
7512            )?;
7513            for token in 0..tokens {
7514                let mut gate = 0.0;
7515                for column in 0..hidden {
7516                    gate += x[token * hidden + column] * gate_weight[column];
7517                }
7518                let gate = sigmoid(gate);
7519                for column in 0..hidden {
7520                    shared_output[token * hidden + column] *= gate;
7521                }
7522            }
7523        }
7524        add_in_place(&mut output, &shared_output);
7525    }
7526    Ok(output)
7527}
7528
7529fn route_experts(
7530    router: &memra_gguf::model_plan::RouterPlan,
7531    logits: &[f32],
7532    bias: Option<&[f32]>,
7533    selected: usize,
7534    forced_indices: Option<&[usize]>,
7535    layer: u32,
7536) -> Result<Vec<(usize, f32)>, ReferenceError> {
7537    use memra_gguf::model_plan::{RouterPlan, RouterScorePlan};
7538
7539    let mut weights = match router {
7540        RouterPlan::Softmax => {
7541            let mut probabilities = logits.to_vec();
7542            softmax_in_place(&mut probabilities);
7543            probabilities
7544        }
7545        RouterPlan::Sigmoid { .. } => logits.iter().map(|&value| sigmoid(value)).collect(),
7546        RouterPlan::SqrtSoftplus { .. } => {
7547            logits.iter().map(|&value| softplus(value).sqrt()).collect()
7548        }
7549        RouterPlan::TokenIdHash { score, .. } => match score {
7550            RouterScorePlan::Softmax => {
7551                let mut probabilities = logits.to_vec();
7552                softmax_in_place(&mut probabilities);
7553                probabilities
7554            }
7555            RouterScorePlan::Sigmoid => logits.iter().map(|&value| sigmoid(value)).collect(),
7556            RouterScorePlan::SqrtSoftplus => {
7557                logits.iter().map(|&value| softplus(value).sqrt()).collect()
7558            }
7559        },
7560    };
7561    let selection_scores: Vec<f32> = weights
7562        .iter()
7563        .enumerate()
7564        .map(|(index, &weight)| weight + bias.map_or(0.0, |bias| bias[index]))
7565        .collect();
7566    let indices = if let RouterPlan::TokenIdHash { .. } = router {
7567        let Some(forced) = forced_indices else {
7568            return Err(ReferenceError::InvalidPlan {
7569                layer: Some(layer),
7570                reason: "token-id hash router requires a token-to-expert row",
7571            });
7572        };
7573        if forced.len() != selected {
7574            return Err(ReferenceError::InvalidPlan {
7575                layer: Some(layer),
7576                reason: "token-id expert row width does not match MoE top-k",
7577            });
7578        }
7579        let mut seen = std::collections::BTreeSet::new();
7580        for &index in forced {
7581            if index >= logits.len() || !seen.insert(index) {
7582                return Err(ReferenceError::InvalidPlan {
7583                    layer: Some(layer),
7584                    reason: "token-id expert row contains an out-of-range or duplicate expert",
7585                });
7586            }
7587        }
7588        forced.to_vec()
7589    } else {
7590        if forced_indices.is_some() {
7591            return Err(ReferenceError::InvalidPlan {
7592                layer: Some(layer),
7593                reason: "score-selected router received forced expert indices",
7594            });
7595        }
7596        let mut indices: Vec<usize> = (0..logits.len()).collect();
7597        indices.sort_by(|&left, &right| {
7598            selection_scores[right]
7599                .total_cmp(&selection_scores[left])
7600                .then(left.cmp(&right))
7601        });
7602        indices.truncate(selected);
7603        indices
7604    };
7605    let (normalize, scaling) = match router {
7606        RouterPlan::Softmax => (true, 1.0),
7607        RouterPlan::Sigmoid {
7608            normalize_selected,
7609            scaling_factor,
7610            ..
7611        }
7612        | RouterPlan::SqrtSoftplus {
7613            normalize_selected,
7614            scaling_factor,
7615            ..
7616        } => (*normalize_selected, *scaling_factor),
7617        RouterPlan::TokenIdHash {
7618            normalize_selected,
7619            scaling_factor,
7620            ..
7621        } => (*normalize_selected, *scaling_factor),
7622    };
7623    if normalize {
7624        let denominator = indices
7625            .iter()
7626            .map(|&index| weights[index])
7627            .sum::<f32>()
7628            .max(if matches!(router, RouterPlan::Softmax) {
7629                6.103_515_6e-5
7630            } else {
7631                1e-20
7632            });
7633        for weight in &mut weights {
7634            *weight = *weight / denominator * scaling;
7635        }
7636    } else {
7637        for weight in &mut weights {
7638            *weight *= scaling;
7639        }
7640    }
7641    Ok(indices
7642        .into_iter()
7643        .map(|index| (index, weights[index]))
7644        .collect())
7645}
7646
7647fn router_has_selection_bias(router: &memra_gguf::model_plan::RouterPlan) -> bool {
7648    matches!(
7649        router,
7650        memra_gguf::model_plan::RouterPlan::Sigmoid {
7651            selection_bias: true,
7652            ..
7653        } | memra_gguf::model_plan::RouterPlan::SqrtSoftplus {
7654            selection_bias: true,
7655            ..
7656        }
7657    )
7658}
7659
7660fn activate_pair(
7661    activation: &ActivationPlan,
7662    gate: f32,
7663    up: f32,
7664    layer: u32,
7665) -> Result<f32, ReferenceError> {
7666    Ok(match activation {
7667        ActivationPlan::Silu => silu(gate) * up,
7668        ActivationPlan::GeluTanh => gelu_tanh(gate) * up,
7669        ActivationPlan::SwiGluOai { alpha, limit } => {
7670            (gate * sigmoid(*alpha * gate)).min(*limit) * up.clamp(-*limit, *limit)
7671        }
7672        ActivationPlan::SwiGluClamped { limit } => {
7673            silu(gate).min(*limit) * up.clamp(-*limit, *limit)
7674        }
7675        // glm5_next: the gate clamp is PRE-silu and one-sided (no lower bound).
7676        ActivationPlan::SwiGluPreClamped { limit } => {
7677            silu(gate.min(*limit)) * up.clamp(-*limit, *limit)
7678        }
7679        ActivationPlan::Named(_) => {
7680            return Err(ReferenceError::UnsupportedOperation {
7681                layer: Some(layer),
7682                operation: "named MLP activation",
7683            });
7684        }
7685    })
7686}
7687
7688fn tensor<'a>(
7689    weights: &'a ReferenceWeights,
7690    id: &TensorId,
7691    expected: &[usize],
7692) -> Result<&'a [f32], ReferenceError> {
7693    let tensor = weights
7694        .get(id)
7695        .ok_or_else(|| ReferenceError::MissingTensor(id.clone()))?;
7696    tensor_checked(id, tensor, expected)
7697}
7698
7699fn tensor_checked<'a>(
7700    id: &TensorId,
7701    tensor: &'a ReferenceTensor,
7702    expected: &[usize],
7703) -> Result<&'a [f32], ReferenceError> {
7704    if tensor.shape != expected {
7705        return Err(ReferenceError::TensorShape {
7706            id: Some(id.clone()),
7707            expected: expected.to_vec(),
7708            actual_elements: tensor.data.len(),
7709        });
7710    }
7711    Ok(&tensor.data)
7712}
7713
7714fn layer_id(layer: u32, tensor: LayerTensor) -> TensorId {
7715    TensorId::Layer {
7716        index: layer,
7717        tensor,
7718    }
7719}
7720
7721fn linear(x: &[f32], weight: &[f32], rows: usize, input: usize, output: usize) -> Vec<f32> {
7722    let mut result = vec![0.0; rows * output];
7723    for row in 0..rows {
7724        for out in 0..output {
7725            let mut sum = 0.0;
7726            for inner in 0..input {
7727                sum += x[row * input + inner] * weight[out * input + inner];
7728            }
7729            result[row * output + out] = sum;
7730        }
7731    }
7732    result
7733}
7734
7735fn rms_norm(x: &[f32], rows: usize, width: usize, weight: &[f32], epsilon: f32) -> Vec<f32> {
7736    let mut result = vec![0.0; x.len()];
7737    for row in 0..rows {
7738        let input = &x[row * width..(row + 1) * width];
7739        let mean_square = input.iter().map(|value| value * value).sum::<f32>() / width as f32;
7740        let inverse = 1.0 / (mean_square + epsilon).sqrt();
7741        for index in 0..width {
7742            result[row * width + index] = input[index] * inverse * weight[index];
7743        }
7744    }
7745    result
7746}
7747
7748/// LayerNorm WITH bias (indexer k_norm). The epsilon is nn.LayerNorm's default and
7749/// does NOT track rms_norm_eps.
7750fn layer_norm(x: &[f32], rows: usize, width: usize, weight: &[f32], bias: &[f32]) -> Vec<f32> {
7751    const EPSILON: f32 = 1e-5;
7752    let mut result = vec![0.0; x.len()];
7753    for row in 0..rows {
7754        let input = &x[row * width..(row + 1) * width];
7755        let mean = input.iter().sum::<f32>() / width as f32;
7756        let variance = input
7757            .iter()
7758            .map(|value| (value - mean) * (value - mean))
7759            .sum::<f32>()
7760            / width as f32;
7761        let inverse = 1.0 / (variance + EPSILON).sqrt();
7762        for index in 0..width {
7763            result[row * width + index] =
7764                (input[index] - mean) * inverse * weight[index] + bias[index];
7765        }
7766    }
7767    result
7768}
7769
7770fn l2_normalize_rows(values: &mut [f32], rows: usize, width: usize, epsilon: f32) {
7771    for row in 0..rows {
7772        let offset = row * width;
7773        let sum = values[offset..offset + width]
7774            .iter()
7775            .map(|value| value * value)
7776            .sum::<f32>();
7777        let inverse = 1.0 / (sum + epsilon).sqrt();
7778        for value in &mut values[offset..offset + width] {
7779            *value *= inverse;
7780        }
7781    }
7782}
7783
7784fn apply_optional_head_norm(
7785    weights: &ReferenceWeights,
7786    id: TensorId,
7787    values: &mut [f32],
7788    rows: usize,
7789    width: usize,
7790    presence: memra_gguf::model_plan::TensorPresence,
7791    epsilon: f32,
7792) -> Result<(), ReferenceError> {
7793    let Some(weight) = weights.get(&id) else {
7794        return if presence == memra_gguf::model_plan::TensorPresence::Required {
7795            Err(ReferenceError::MissingTensor(id))
7796        } else {
7797            Ok(())
7798        };
7799    };
7800    let normalized = rms_norm(
7801        values,
7802        rows,
7803        width,
7804        tensor_checked(&id, weight, &[width])?,
7805        epsilon,
7806    );
7807    values.copy_from_slice(&normalized);
7808    Ok(())
7809}
7810
7811/// Per-dim frequency divisors + the cos/sin attention scale (YaRN mscale; 1.0 for every
7812/// other factor kind — an exact multiplicative identity).
7813fn rope_factor_values(
7814    plan: &memra_gguf::model_plan::RopePlan,
7815    weights: &ReferenceWeights,
7816) -> Result<(Option<Vec<f32>>, f32), ReferenceError> {
7817    use memra_gguf::model_plan::RopeFactors;
7818
7819    let width = plan.dimensions as usize / 2;
7820    Ok(match plan.factors {
7821        RopeFactors::None => (None, 1.0),
7822        RopeFactors::PartialRotary { factor } => {
7823            let keep = (width as f32 * factor.clamp(0.0, 1.0)).round() as usize;
7824            (
7825                Some(
7826                    (0..width)
7827                        .map(|index| if index < keep { 1.0 } else { 1.0e30 })
7828                        .collect(),
7829                ),
7830                1.0,
7831            )
7832        }
7833        RopeFactors::Checkpoint => {
7834            let tensor = weights
7835                .get(&TensorId::RopeFactors)
7836                .ok_or(ReferenceError::MissingTensor(TensorId::RopeFactors))?;
7837            if tensor.shape.len() != 1 || tensor.data.len() < width {
7838                return Err(ReferenceError::TensorShape {
7839                    id: Some(TensorId::RopeFactors),
7840                    expected: vec![width],
7841                    actual_elements: tensor.data.len(),
7842                });
7843            }
7844            (Some(tensor.data[..width].to_vec()), 1.0)
7845        }
7846        // YaRN on full attention (qwen4_exp long-context lane): the transformers-twin
7847        // frequency divisors + the derived attention factor on cos/sin. The divisor table
7848        // shares the Checkpoint-factors convention, so every consumer below (QSA q/k AND
7849        // the indexer's q/pooled-k rope) rides the same path.
7850        RopeFactors::Yarn {
7851            factor,
7852            original_context,
7853            beta_fast,
7854            beta_slow,
7855        } => (
7856            Some(memra_gguf::model_plan::yarn_frequency_divisors(
7857                plan.dimensions,
7858                plan.base,
7859                factor,
7860                original_context,
7861                beta_fast,
7862                beta_slow,
7863            )),
7864            memra_gguf::model_plan::yarn_attention_factor(factor),
7865        ),
7866    })
7867}
7868
7869#[allow(clippy::too_many_arguments)]
7870fn apply_rope(
7871    values: &mut [f32],
7872    tokens: usize,
7873    heads: usize,
7874    head_dim: usize,
7875    dimensions: usize,
7876    base: f32,
7877    factors: Option<&[f32]>,
7878    mscale: f32,
7879) {
7880    for token in 0..tokens {
7881        apply_rope_at_position(
7882            &mut values[token * heads * head_dim..(token + 1) * heads * head_dim],
7883            heads,
7884            head_dim,
7885            dimensions,
7886            base,
7887            factors,
7888            mscale,
7889            token,
7890        );
7891    }
7892}
7893
7894/// One row of NeoX split-half rope at an EXPLICIT position — the QSA indexer rotates
7895/// pooled block keys at the block-start position, not their row index. `mscale` is the
7896/// YaRN attention factor on cos/sin (transformers `attention_scaling`; 1.0 elsewhere —
7897/// an exact multiplicative identity, so the non-yarn arms are byte-unchanged).
7898#[allow(clippy::too_many_arguments)]
7899fn apply_rope_at_position(
7900    values: &mut [f32],
7901    heads: usize,
7902    head_dim: usize,
7903    dimensions: usize,
7904    base: f32,
7905    factors: Option<&[f32]>,
7906    mscale: f32,
7907    position: usize,
7908) {
7909    let dimensions = dimensions.min(head_dim) / 2 * 2;
7910    let half = dimensions / 2;
7911    for head in 0..heads {
7912        let offset = head * head_dim;
7913        for index in 0..half {
7914            let factor = factors.map_or(1.0, |factors| factors[index]);
7915            let frequency = base.powf(-2.0 * index as f32 / dimensions as f32) / factor;
7916            let angle = position as f32 * frequency;
7917            let (sin, cos) = angle.sin_cos();
7918            let (sin, cos) = (sin * mscale, cos * mscale);
7919            let first = values[offset + index];
7920            let second = values[offset + index + half];
7921            values[offset + index] = first * cos - second * sin;
7922            values[offset + index + half] = first * sin + second * cos;
7923        }
7924    }
7925}
7926
7927fn softmax_in_place(values: &mut [f32]) {
7928    let max = values.iter().copied().fold(f32::NEG_INFINITY, f32::max);
7929    let mut sum = 0.0;
7930    for value in values.iter_mut() {
7931        *value = (*value - max).exp();
7932        sum += *value;
7933    }
7934    for value in values {
7935        *value /= sum;
7936    }
7937}
7938
7939fn add_in_place(target: &mut [f32], addend: &[f32]) {
7940    for (target, addend) in target.iter_mut().zip(addend) {
7941        *target += addend;
7942    }
7943}
7944
7945fn sigmoid(value: f32) -> f32 {
7946    1.0 / (1.0 + (-value).exp())
7947}
7948
7949fn silu(value: f32) -> f32 {
7950    value * sigmoid(value)
7951}
7952
7953fn softplus(value: f32) -> f32 {
7954    if value > 20.0 {
7955        value
7956    } else {
7957        value.exp().ln_1p()
7958    }
7959}
7960
7961fn gelu_tanh(value: f32) -> f32 {
7962    0.5 * value * (1.0 + (0.797_884_6 * (value + 0.044_715 * value * value * value)).tanh())
7963}
7964
7965/// Exact-erf GELU (torch `nn.GELU()` default, used by the glm5_next vision merger — NOT
7966/// the tanh approximation). erf via Abramowitz & Stegun 7.1.26 in f64 (max abs error
7967/// 1.5e-7, below f32 resolution at these magnitudes).
7968fn gelu_erf(value: f32) -> f32 {
7969    let x = value as f64 / std::f64::consts::SQRT_2;
7970    let sign = if x < 0.0 { -1.0 } else { 1.0 };
7971    let x = x.abs();
7972    let t = 1.0 / (1.0 + 0.327_591_1 * x);
7973    let poly = t
7974        * (0.254_829_592
7975            + t * (-0.284_496_736
7976                + t * (1.421_413_741 + t * (-1.453_152_027 + t * 1.061_405_429))));
7977    let erf = sign * (1.0 - poly * (-x * x).exp());
7978    (0.5 * value as f64 * (1.0 + erf)) as f32
7979}
7980
7981#[cfg(test)]
7982mod tests {
7983    use super::*;
7984    use memra_gguf::config::{HfConfig, ModelConfig};
7985
7986    fn weight(shape: &[usize], data: &[f32]) -> ReferenceTensor {
7987        ReferenceTensor::new(shape.to_vec(), data.to_vec()).unwrap()
7988    }
7989
7990    #[test]
7991    fn one_token_dense_plan_matches_hand_derived_logits_and_emits_kv_state() {
7992        let config = ModelConfig::from_hf(&HfConfig::parse(
7993            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":2,
7994            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":2,
7995            "intermediate_size":2,"vocab_size":3,"max_position_embeddings":8,
7996            "rms_norm_eps":0.000001}"#,
7997        ));
7998        let plan = ModelPlan::compile(&config).unwrap();
7999        let identity = [1.0, 0.0, 0.0, 1.0];
8000        let zero = [0.0; 4];
8001        let mut weights = ReferenceWeights::new();
8002        weights.insert(
8003            TensorId::TokenEmbedding,
8004            weight(&[3, 2], &[1.0, 0.0, 0.0, 1.0, -1.0, 0.0]),
8005        );
8006        weights.insert(TensorId::OutputNorm, weight(&[2], &[1.0, 1.0]));
8007        for tensor in [LayerTensor::PreAttentionNorm, LayerTensor::PreMlpNorm] {
8008            weights.insert(layer_id(0, tensor), weight(&[2], &[1.0, 1.0]));
8009        }
8010        for tensor in [
8011            LayerTensor::Query,
8012            LayerTensor::Key,
8013            LayerTensor::Value,
8014            LayerTensor::AttentionOutput,
8015        ] {
8016            weights.insert(layer_id(0, tensor), weight(&[2, 2], &identity));
8017        }
8018        for tensor in [
8019            LayerTensor::MlpGate,
8020            LayerTensor::MlpUp,
8021            LayerTensor::MlpDown,
8022        ] {
8023            weights.insert(layer_id(0, tensor), weight(&[2, 2], &zero));
8024        }
8025
8026        let output = execute(&plan, &weights, &[0]).unwrap();
8027        let root_two = 2.0f32.sqrt();
8028        assert_eq!((output.tokens, output.vocab), (1, 3));
8029        assert!((output.logits[0] - root_two).abs() < 2e-5);
8030        assert!(output.logits[1].abs() < 2e-5);
8031        assert!((output.logits[2] + root_two).abs() < 2e-5);
8032        let ReferenceLayerState::Kv {
8033            tokens, key, value, ..
8034        } = &output.state.layers[0]
8035        else {
8036            panic!("expected KV state");
8037        };
8038        assert_eq!(*tokens, 1);
8039        assert_eq!(key.len(), 2);
8040        assert_eq!(value.len(), 2);
8041    }
8042
8043    #[test]
8044    fn hyperconnections_execute_stream_state_and_head_collapse() {
8045        let config = ModelConfig::from_hf(&HfConfig::parse(
8046            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":2,
8047            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":2,
8048            "intermediate_size":2,"vocab_size":3,"max_position_embeddings":8}"#,
8049        ));
8050        let mut plan = ModelPlan::compile(&config).unwrap();
8051        plan.layers[0].residual = ResidualTopology::HyperConnections {
8052            streams: 2,
8053            epsilon: 1e-6,
8054            sinkhorn_iterations: 2,
8055            collapse: HcCollapse::GatedHead,
8056        };
8057        let fixture = deterministic_fixture(&plan).unwrap();
8058        assert_eq!(
8059            fixture.weights[&TensorId::HyperHeadFunction].shape,
8060            vec![2, 4]
8061        );
8062        assert_eq!(
8063            fixture.weights[&layer_id(0, LayerTensor::HyperAttentionFunction)].shape,
8064            vec![8, 4]
8065        );
8066        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8067        assert!(output.logits.iter().all(|value| value.is_finite()));
8068        assert!(matches!(
8069            output.state.layers[0],
8070            ReferenceLayerState::Kv { .. }
8071        ));
8072    }
8073
8074    #[test]
8075    fn generated_tiny_fixture_is_deterministic_and_executable() {
8076        let config = ModelConfig::from_hf(&HfConfig::parse(
8077            r#"{"model_type":"qwen3","num_hidden_layers":2,"hidden_size":8,
8078            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8079            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32}"#,
8080        ));
8081        let plan = ModelPlan::compile(&config).unwrap();
8082        let first = deterministic_fixture(&plan).unwrap();
8083        let second = deterministic_fixture(&plan).unwrap();
8084        assert_eq!(first, second);
8085        let output = execute(&plan, &first.weights, &first.token_ids).unwrap();
8086        assert_eq!(output.logits.len(), first.token_ids.len() * 32);
8087        assert!(output.logits.iter().all(|value| value.is_finite()));
8088    }
8089
8090    #[test]
8091    fn qwen35_fixture_executes_mixed_gdn_and_full_attention_state() {
8092        let config = ModelConfig::from_hf(&HfConfig::parse(
8093            r#"{"model_type":"qwen3_5","num_hidden_layers":4,"hidden_size":8,
8094            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8095            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
8096            "rms_norm_eps":0.000001,"full_attention_interval":2,
8097            "linear_conv_kernel_dim":3,"linear_key_head_dim":4,
8098            "linear_value_head_dim":4,"linear_num_key_heads":1,
8099            "linear_num_value_heads":2}"#,
8100        ));
8101        let plan = ModelPlan::compile(&config).unwrap();
8102        let fixture = deterministic_fixture(&plan).unwrap();
8103        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8104        assert_eq!(output.state.layers.len(), 4);
8105        assert!(matches!(
8106            output.state.layers[0],
8107            ReferenceLayerState::Recurrent { .. }
8108        ));
8109        assert!(matches!(
8110            output.state.layers[1],
8111            ReferenceLayerState::Kv { .. }
8112        ));
8113        assert!(matches!(
8114            output.state.layers[2],
8115            ReferenceLayerState::Recurrent { .. }
8116        ));
8117        assert!(matches!(
8118            output.state.layers[3],
8119            ReferenceLayerState::Kv { .. }
8120        ));
8121        assert!(output.logits.iter().all(|value| value.is_finite()));
8122        assert_eq!(
8123            output.logits[..8]
8124                .iter()
8125                .map(|value| value.to_bits())
8126                .collect::<Vec<_>>(),
8127            vec![
8128                3_182_242_076,
8129                1_053_299_392,
8130                3_199_800_546,
8131                3_198_737_445,
8132                3_180_184_136,
8133                3_187_768_631,
8134                1_057_556_100,
8135                1_035_812_924,
8136            ]
8137        );
8138    }
8139
8140    #[test]
8141    fn router_laws_pin_stable_ties_and_selection_only_bias() {
8142        use memra_gguf::model_plan::{RouterPlan, RouterScorePlan};
8143
8144        assert_eq!(
8145            route_experts(&RouterPlan::Softmax, &[0.0, 0.0, 0.0], None, 2, None, 0,).unwrap(),
8146            vec![(0, 0.5), (1, 0.5)]
8147        );
8148        assert_eq!(
8149            route_experts(
8150                &RouterPlan::Sigmoid {
8151                    normalize_selected: true,
8152                    scaling_factor: 2.0,
8153                    selection_bias: true,
8154                },
8155                &[0.0, 0.0],
8156                Some(&[-1.0, 1.0]),
8157                1,
8158                None,
8159                0,
8160            )
8161            .unwrap(),
8162            vec![(1, 2.0)]
8163        );
8164        assert_eq!(
8165            route_experts(
8166                &RouterPlan::TokenIdHash {
8167                    score: RouterScorePlan::SqrtSoftplus,
8168                    normalize_selected: true,
8169                    scaling_factor: 1.5,
8170                },
8171                &[0.0, 0.0, 0.0],
8172                None,
8173                2,
8174                Some(&[2, 0]),
8175                0,
8176            )
8177            .unwrap(),
8178            vec![(2, 0.75), (0, 0.75)]
8179        );
8180        assert!(matches!(
8181            route_experts(
8182                &RouterPlan::TokenIdHash {
8183                    score: RouterScorePlan::SqrtSoftplus,
8184                    normalize_selected: true,
8185                    scaling_factor: 1.5,
8186                },
8187                &[0.0, 0.0, 0.0],
8188                None,
8189                2,
8190                Some(&[1, 1]),
8191                0,
8192            ),
8193            Err(ReferenceError::InvalidPlan {
8194                reason: "token-id expert row contains an out-of-range or duplicate expert",
8195                ..
8196            })
8197        ));
8198    }
8199
8200    #[test]
8201    fn token_hash_moe_fixture_executes_from_semantic_token_table() {
8202        use memra_gguf::model_plan::{RouterPlan, RouterScorePlan};
8203
8204        let config = ModelConfig::from_hf(&HfConfig::parse(
8205            r#"{"model_type":"qwen3_moe","num_hidden_layers":1,"hidden_size":8,
8206            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8207            "intermediate_size":16,"vocab_size":16,"max_position_embeddings":32,
8208            "num_experts":4,"num_experts_per_tok":2,"moe_intermediate_size":8}"#,
8209        ));
8210        let mut plan = ModelPlan::compile(&config).unwrap();
8211        let MlpPlan::Moe(moe) = &mut plan.layers[0].mlp else {
8212            unreachable!()
8213        };
8214        moe.router = RouterPlan::TokenIdHash {
8215            score: RouterScorePlan::SqrtSoftplus,
8216            normalize_selected: true,
8217            scaling_factor: 1.5,
8218        };
8219        let fixture = deterministic_fixture(&plan).unwrap();
8220        let table_id = layer_id(0, LayerTensor::MoeTokenToExpert);
8221        assert_eq!(fixture.weights[&table_id].shape, vec![16, 2]);
8222        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8223        assert!(output.logits.iter().all(|value| value.is_finite()));
8224
8225        let mut alternate = fixture.weights.clone();
8226        alternate.get_mut(&table_id).unwrap().data.fill(3.0);
8227        for row in alternate
8228            .get_mut(&table_id)
8229            .unwrap()
8230            .data
8231            .chunks_exact_mut(2)
8232        {
8233            row[1] = 2.0;
8234        }
8235        let alternate = execute(&plan, &alternate, &fixture.token_ids).unwrap();
8236        assert_ne!(output.logits, alternate.logits);
8237    }
8238
8239    #[test]
8240    fn qwen3_moe_fixture_executes_routed_and_shared_branches() {
8241        let config = ModelConfig::from_hf(&HfConfig::parse(
8242            r#"{"model_type":"qwen3_moe","num_hidden_layers":2,"hidden_size":8,
8243            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8244            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
8245            "num_experts":4,"num_experts_per_tok":2,"moe_intermediate_size":8,
8246            "shared_expert_intermediate_size":8}"#,
8247        ));
8248        let plan = ModelPlan::compile(&config).unwrap();
8249        let fixture = deterministic_fixture(&plan).unwrap();
8250        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8251        assert!(output.logits.iter().all(|value| value.is_finite()));
8252        assert_eq!(
8253            output.logits[..8]
8254                .iter()
8255                .map(|value| value.to_bits())
8256                .collect::<Vec<_>>(),
8257            vec![
8258                3_205_834_204,
8259                1_034_800_117,
8260                1_053_917_366,
8261                3_190_866_844,
8262                984_171_488,
8263                3_182_514_784,
8264                3_154_736_064,
8265                3_175_624_690,
8266            ]
8267        );
8268    }
8269
8270    #[test]
8271    fn sliding_window_limits_attention_and_trims_reference_state() {
8272        let config = ModelConfig::from_hf(&HfConfig::parse(
8273            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":8,
8274            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8275            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32}"#,
8276        ));
8277        let mut plan = ModelPlan::compile(&config).unwrap();
8278        let AttentionPlan::Full(attention) = plan.layers[0].attention.clone() else {
8279            unreachable!()
8280        };
8281        plan.layers[0].attention = AttentionPlan::SlidingWindow {
8282            attention,
8283            window: 2,
8284        };
8285        let fixture = deterministic_fixture(&plan).unwrap();
8286        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8287        let ReferenceLayerState::Kv { tokens, window, .. } = output.state.layers[0] else {
8288            panic!("expected sliding KV state");
8289        };
8290        assert_eq!(tokens, 2);
8291        assert_eq!(window, Some(2));
8292    }
8293
8294    #[test]
8295    fn mla_fixture_emits_latent_state_and_sparse_overflow_refuses() {
8296        use memra_gguf::model_plan::{
8297            MlaAttentionPlan, RopeFactors, RopePlan, SparseIndexPlan, StatePlan,
8298        };
8299
8300        let config = ModelConfig::from_hf(&HfConfig::parse(
8301            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":8,
8302            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8303            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32}"#,
8304        ));
8305        let mut plan = ModelPlan::compile(&config).unwrap();
8306        let mla = MlaAttentionPlan::LatentKv {
8307            query_heads: 2,
8308            q_lora_rank: 4,
8309            kv_lora_rank: 4,
8310            qk_head_dim: 4,
8311            rope_head_dim: 2,
8312            value_head_dim: 4,
8313            rope: RopePlan {
8314                dimensions: 2,
8315                base: 10_000.0,
8316                factors: RopeFactors::None,
8317            },
8318            sparse_index: SparseIndexPlan::None,
8319        };
8320        plan.layers[0].attention = AttentionPlan::Mla(mla.clone());
8321        plan.layers[0].state = StatePlan::LatentKvCache {
8322            width: 6,
8323            index_width: 0,
8324        };
8325        let fixture = deterministic_fixture(&plan).unwrap();
8326        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8327        let ReferenceLayerState::LatentKv { tokens, width, .. } = output.state.layers[0] else {
8328            panic!("expected latent KV state");
8329        };
8330        assert_eq!((tokens, width), (3, 6));
8331        assert_eq!(
8332            output.logits[..4]
8333                .iter()
8334                .map(|value| value.to_bits())
8335                .collect::<Vec<_>>(),
8336            vec![1_035_177_220, 1_055_447_641, 3_201_478_680, 3_199_508_856]
8337        );
8338
8339        let MlaAttentionPlan::LatentKv {
8340            query_heads,
8341            q_lora_rank,
8342            kv_lora_rank,
8343            qk_head_dim,
8344            rope_head_dim,
8345            value_head_dim,
8346            rope,
8347            ..
8348        } = mla
8349        else {
8350            unreachable!()
8351        };
8352        plan.layers[0].attention = AttentionPlan::Mla(MlaAttentionPlan::LatentKv {
8353            query_heads,
8354            q_lora_rank,
8355            kv_lora_rank,
8356            qk_head_dim,
8357            rope_head_dim,
8358            value_head_dim,
8359            rope,
8360            sparse_index: SparseIndexPlan::Own {
8361                heads: 1,
8362                head_dim: 2,
8363                top_k: 2,
8364                kpool: None,
8365            },
8366        });
8367        let error = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap_err();
8368        assert!(matches!(
8369            error,
8370            ReferenceError::UnsupportedOperation {
8371                operation: "sparse MLA selection beyond full-selection equivalence",
8372                ..
8373            }
8374        ));
8375    }
8376
8377    #[test]
8378    fn compressed_mla_executes_window_compressor_indexer_and_grouped_output() {
8379        use memra_gguf::model_plan::{
8380            KvCompressorPlan, MlaAttentionPlan, RopeFactors, RopePlan, SparseIndexPlan, StatePlan,
8381        };
8382
8383        let config = ModelConfig::from_hf(&HfConfig::parse(
8384            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":128,
8385            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":64,
8386            "intermediate_size":256,"vocab_size":32,"max_position_embeddings":64,
8387            "rms_norm_eps":0.000001}"#,
8388        ));
8389        let mut plan = ModelPlan::compile(&config).unwrap();
8390        plan.layers[0].attention = AttentionPlan::Mla(MlaAttentionPlan::CompressedKv {
8391            query_heads: 2,
8392            q_lora_rank: 64,
8393            latent_head_dim: 128,
8394            rope_head_dim: 64,
8395            output_lora_rank: 64,
8396            output_groups: 1,
8397            window: 4,
8398            rope: RopePlan {
8399                dimensions: 64,
8400                base: 160_000.0,
8401                factors: RopeFactors::Yarn {
8402                    factor: 2.0,
8403                    original_context: 32,
8404                    beta_fast: 32.0,
8405                    beta_slow: 1.0,
8406                },
8407            },
8408            compressor: Some(KvCompressorPlan {
8409                ratio: 4,
8410                latent_dim: 256,
8411            }),
8412            sparse_index: SparseIndexPlan::Own {
8413                heads: 2,
8414                head_dim: 128,
8415                top_k: 2,
8416                kpool: None,
8417            },
8418        });
8419        plan.layers[0].state = StatePlan::CompressedAttention {
8420            window: 4,
8421            head_dim: 128,
8422            compressor_ratio: Some(4),
8423            sparse_top_k: Some(2),
8424        };
8425        let fixture = deterministic_fixture(&plan).unwrap();
8426        let output = execute(&plan, &fixture.weights, &[1, 2, 3, 4]).unwrap();
8427        let ReferenceLayerState::CompressedAttention {
8428            tokens,
8429            width,
8430            window,
8431            compressed_tokens,
8432            ..
8433        } = output.state.layers[0]
8434        else {
8435            panic!("expected compressed attention state")
8436        };
8437        assert_eq!((tokens, width, window, compressed_tokens), (5, 128, 4, 1));
8438        assert!(output.logits.iter().all(|value| value.is_finite()));
8439    }
8440
8441    #[test]
8442    fn dsv4_shaped_trunk_executes_one_canonical_plan() {
8443        let config = ModelConfig::from_hf(&HfConfig::parse(
8444            r#"{"model_type":"deepseek_v4","num_hidden_layers":2,"hidden_size":128,
8445            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":128,
8446            "intermediate_size":256,"vocab_size":128,"max_position_embeddings":1024,
8447            "rms_norm_eps":0.000001,"rope_theta":10000,"n_routed_experts":4,
8448            "n_shared_experts":1,"num_experts_per_tok":2,"moe_intermediate_size":128,
8449            "norm_topk_prob":true,"num_hash_layers":1,"num_nextn_predict_layers":1,
8450            "scoring_func":"sqrtsoftplus","topk_method":"noaux_tc",
8451            "routed_scaling_factor":1.5,"hc_eps":0.000001,"hc_mult":2,
8452            "hc_sinkhorn_iters":4,"q_lora_rank":128,"qk_rope_head_dim":64,
8453            "o_lora_rank":128,"o_groups":1,"index_n_heads":1,"index_head_dim":128,
8454            "index_topk":16,"compress_ratios":[0,4,0],"compress_rope_theta":160000,
8455            "sliding_window":128,"swiglu_limit":10.0,
8456            "rope_scaling":{"factor":4,"beta_fast":32,"beta_slow":1,
8457            "original_max_position_embeddings":1024}}"#,
8458        ));
8459        let mut plan = ModelPlan::compile(&config).unwrap();
8460        assert_eq!(plan.layers.len(), 2);
8461        plan.mtp_blocks.clear();
8462        let fixture = deterministic_fixture(&plan).unwrap();
8463        let output = execute(&plan, &fixture.weights, &[1, 2, 3, 4]).unwrap();
8464        assert_eq!(output.state.layers.len(), 2);
8465        assert!(
8466            output
8467                .state
8468                .layers
8469                .iter()
8470                .all(|state| matches!(state, ReferenceLayerState::CompressedAttention { .. }))
8471        );
8472        assert!(
8473            fixture
8474                .weights
8475                .contains_key(&layer_id(0, LayerTensor::MoeTokenToExpert))
8476        );
8477        assert!(
8478            fixture
8479                .weights
8480                .contains_key(&layer_id(1, LayerTensor::MoeRouterBias))
8481        );
8482        assert!(output.logits.iter().all(|value| value.is_finite()));
8483    }
8484
8485    #[test]
8486    fn dspark_executes_trunk_tap_ring_blocks_markov_and_confidence() {
8487        use memra_gguf::model_plan::{DrafterPlan, DsparkPlan};
8488
8489        let config = ModelConfig::from_hf(&HfConfig::parse(
8490            r#"{"model_type":"deepseek_v4","num_hidden_layers":2,"hidden_size":128,
8491            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":128,
8492            "intermediate_size":256,"vocab_size":128,"max_position_embeddings":1024,
8493            "rms_norm_eps":0.000001,"rope_theta":10000,"n_routed_experts":4,
8494            "n_shared_experts":1,"num_experts_per_tok":2,"moe_intermediate_size":128,
8495            "norm_topk_prob":true,"num_hash_layers":1,"num_nextn_predict_layers":1,
8496            "scoring_func":"sqrtsoftplus","topk_method":"noaux_tc",
8497            "routed_scaling_factor":1.5,"hc_eps":0.000001,"hc_mult":2,
8498            "hc_sinkhorn_iters":4,"q_lora_rank":128,"qk_rope_head_dim":64,
8499            "o_lora_rank":128,"o_groups":1,"index_n_heads":1,"index_head_dim":128,
8500            "index_topk":16,"compress_ratios":[0,4,0],"compress_rope_theta":160000,
8501            "sliding_window":128,"swiglu_limit":10.0,
8502            "rope_scaling":{"factor":4,"beta_fast":32,"beta_slow":1,
8503            "original_max_position_embeddings":1024}}"#,
8504        ));
8505        let mut plan = ModelPlan::compile(&config).unwrap();
8506        let block = plan.mtp_blocks.remove(0).layer;
8507        plan.drafter = Some(DrafterPlan::Dspark(DsparkPlan {
8508            block_size: 3,
8509            noise_token_id: 31,
8510            target_layer_ids: vec![1],
8511            markov_rank: 8,
8512            blocks: vec![block],
8513        }));
8514        let fixture = deterministic_fixture(&plan).unwrap();
8515        let output = execute(&plan, &fixture.weights, &[1, 2, 3, 4]).unwrap();
8516        let draft = output.draft.expect("DSpark output");
8517        assert_eq!(draft.input_token, 4);
8518        assert_eq!(draft.output_ids.len(), 4);
8519        assert_eq!(draft.confidence.len(), 3);
8520        assert_eq!(draft.logits.len(), 3 * 128);
8521        assert!(draft.logits.iter().all(|value| value.is_finite()));
8522        assert!(draft.confidence.iter().all(|value| value.is_finite()));
8523    }
8524
8525    #[test]
8526    fn gemma4_vision_executes_patch_rope_pool_standardize_and_projection() {
8527        let config = ModelConfig::from_hf(&HfConfig::parse(
8528            r#"{"model_type":"gemma4","image_token_id":31,"vision_soft_tokens_per_image":1,
8529            "text_config":{"model_type":"gemma4_text",
8530            "num_hidden_layers":2,"hidden_size":8,"num_attention_heads":2,
8531            "num_key_value_heads":1,"num_global_key_value_heads":1,"head_dim":4,
8532            "global_head_dim":4,"intermediate_size":16,"vocab_size":32,
8533            "max_position_embeddings":64,"rms_norm_eps":0.000001,"sliding_window":8,
8534            "layer_types":["sliding_attention","full_attention"],
8535            "rope_parameters":{"full_attention":{"rope_theta":10000,
8536            "partial_rotary_factor":0.5},"sliding_attention":{"rope_theta":10000}}},
8537            "vision_config":{"hidden_size":8,"intermediate_size":16,
8538            "num_hidden_layers":2,"num_attention_heads":2,"num_key_value_heads":1,
8539            "head_dim":4,"max_position_embeddings":64,"patch_size":2,
8540            "position_embedding_size":16,"pooling_kernel_size":2,
8541            "rms_norm_eps":0.000001,"standardize":true,"use_clipped_linears":false,
8542            "hidden_activation":"gelu_pytorch_tanh","rope_parameters":{"rope_theta":100}}}"#,
8543        ));
8544        let plan = ModelPlan::compile(&config).unwrap();
8545        let fixture = deterministic_fixture(&plan).unwrap();
8546        let input = fixture.vision.as_ref().expect("vision fixture");
8547        let first = execute_vision(&plan, &fixture.weights, input).unwrap();
8548        let second = execute_vision(&plan, &fixture.weights, input).unwrap();
8549        assert_eq!(first, second);
8550        assert_eq!((first.patch_count, first.output_tokens), (4, 1));
8551        assert_eq!((first.hidden_size, first.projection_size), (8, 8));
8552        assert_eq!(first.encoder_hidden.len(), 4 * 8);
8553        assert_eq!(first.pooled_hidden.len(), 8);
8554        assert_eq!(first.projected_hidden.len(), 8);
8555        assert!(first.projected_hidden.iter().all(|value| value.is_finite()));
8556        let multimodal = execute_multimodal(&plan, &fixture.weights, &[1, 31, 2], input).unwrap();
8557        let text_only = execute(&plan, &fixture.weights, &[1, 31, 2]).unwrap();
8558        assert_eq!(multimodal.vision, first);
8559        assert_ne!(multimodal.language.logits, text_only.logits);
8560        assert!(
8561            plan.operations()
8562                .contains(&memra_gguf::model_plan::OperationKind::VisionTokenInjection)
8563        );
8564    }
8565
8566    #[test]
8567    fn gemma4_parallel_moe_executes_shared_routed_and_scaled_residual_branches() {
8568        let config = ModelConfig::from_hf(&HfConfig::parse(
8569            r#"{"model_type":"gemma4","text_config":{"model_type":"gemma4_text",
8570            "num_hidden_layers":2,"hidden_size":8,"num_attention_heads":2,
8571            "num_key_value_heads":1,"num_global_key_value_heads":1,"head_dim":4,
8572            "global_head_dim":4,"intermediate_size":16,"moe_intermediate_size":8,
8573            "num_experts":4,"top_k_experts":2,"vocab_size":32,
8574            "max_position_embeddings":64,"rms_norm_eps":0.000001,"sliding_window":8,
8575            "layer_types":["sliding_attention","full_attention"],
8576            "rope_parameters":{"full_attention":{"rope_theta":10000,
8577            "partial_rotary_factor":0.5},"sliding_attention":{"rope_theta":10000}}}}"#,
8578        ));
8579        let plan = ModelPlan::compile(&config).unwrap();
8580        let MlpPlan::Moe(moe) = &plan.layers[0].mlp else {
8581            panic!("expected Gemma MoE")
8582        };
8583        assert_eq!(moe.experts_per_token, 2);
8584        assert_eq!(moe.shared.as_ref().unwrap().intermediate_size, 16);
8585        assert!(matches!(
8586            plan.layers[0].residual,
8587            ResidualTopology::Gemma {
8588                parallel_moe: Some(_),
8589                ..
8590            }
8591        ));
8592        let fixture = deterministic_fixture(&plan).unwrap();
8593        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8594        assert!(output.logits.iter().all(|value| value.is_finite()));
8595        assert!(
8596            plan.operations()
8597                .contains(&memra_gguf::model_plan::OperationKind::GemmaParallelMoeResidual)
8598        );
8599    }
8600
8601    #[test]
8602    fn embedded_mtp_executes_typed_fusion_block_and_fallback_head() {
8603        let config = ModelConfig::from_hf(&HfConfig::parse(
8604            r#"{"model_type":"qwen3_5","num_hidden_layers":2,
8605            "num_nextn_predict_layers":1,"hidden_size":8,
8606            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8607            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
8608            "rms_norm_eps":0.000001,"full_attention_interval":2,
8609            "linear_conv_kernel_dim":3,"linear_key_head_dim":4,
8610            "linear_value_head_dim":4,"linear_num_key_heads":1,
8611            "linear_num_value_heads":2}"#,
8612        ));
8613        let plan = ModelPlan::compile(&config).unwrap();
8614        assert_eq!(plan.mtp_blocks.len(), 1);
8615        let fixture = deterministic_fixture(&plan).unwrap();
8616        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8617        assert_eq!(output.mtp.len(), 1);
8618        assert_eq!(output.mtp[0].depth, 0);
8619        assert_eq!(output.mtp[0].hidden.len(), fixture.token_ids.len() * 8);
8620        assert_eq!(output.mtp[0].logits.len(), fixture.token_ids.len() * 32);
8621        assert!(output.mtp[0].logits.iter().all(|value| value.is_finite()));
8622        assert_eq!(
8623            output.mtp[0].logits[..4]
8624                .iter()
8625                .map(|value| value.to_bits())
8626                .collect::<Vec<_>>(),
8627            vec![1_042_962_358, 1_044_718_512, 3_171_782_004, 3_189_261_409]
8628        );
8629    }
8630
8631    #[test]
8632    fn multi_depth_mtp_threads_hidden_through_every_typed_block() {
8633        let config = ModelConfig::from_hf(&HfConfig::parse(
8634            r#"{"model_type":"qwen3_5","num_hidden_layers":2,
8635            "num_nextn_predict_layers":2,"hidden_size":8,
8636            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8637            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
8638            "rms_norm_eps":0.000001,"full_attention_interval":2,
8639            "linear_conv_kernel_dim":3,"linear_key_head_dim":4,
8640            "linear_value_head_dim":4,"linear_num_key_heads":1,
8641            "linear_num_value_heads":2}"#,
8642        ));
8643        let plan = ModelPlan::compile(&config).unwrap();
8644        assert_eq!(plan.mtp_blocks.len(), 2);
8645        let fixture = deterministic_fixture(&plan).unwrap();
8646        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8647        assert_eq!(
8648            output
8649                .mtp
8650                .iter()
8651                .map(|block| block.depth)
8652                .collect::<Vec<_>>(),
8653            vec![0, 1]
8654        );
8655        assert!(
8656            output
8657                .mtp
8658                .iter()
8659                .flat_map(|block| &block.logits)
8660                .all(|value| value.is_finite())
8661        );
8662        assert_ne!(output.mtp[0].hidden, output.mtp[1].hidden);
8663    }
8664
8665    #[test]
8666    fn rope_uses_neox_split_half_pairs() {
8667        use memra_gguf::model_plan::{RopeFactors, RopePlan};
8668
8669        let mut values = vec![1.0, 2.0, 3.0, 4.0];
8670        apply_rope(&mut values, 1, 1, 4, 4, 10_000.0, None, 1.0);
8671        // Position zero is deliberately unchanged.
8672        assert_eq!(values, vec![1.0, 2.0, 3.0, 4.0]);
8673
8674        let mut values = vec![0.0; 8];
8675        values[4..].copy_from_slice(&[1.0, 2.0, 3.0, 4.0]);
8676        apply_rope(&mut values, 2, 1, 4, 4, 10_000.0, None, 1.0);
8677        let (sin0, cos0) = 1.0f32.sin_cos();
8678        let (sin1, cos1) = 0.01f32.sin_cos();
8679        let row = &values[4..];
8680        assert!((row[0] - (cos0 - 3.0 * sin0)).abs() < 1e-6);
8681        assert!((row[2] - (sin0 + 3.0 * cos0)).abs() < 1e-6);
8682        assert!((row[1] - (2.0 * cos1 - 4.0 * sin1)).abs() < 1e-6);
8683        assert!((row[3] - (2.0 * sin1 + 4.0 * cos1)).abs() < 1e-6);
8684        assert_eq!(
8685            rope_factor_values(
8686                &RopePlan {
8687                    dimensions: 4,
8688                    base: 10_000.0,
8689                    factors: RopeFactors::PartialRotary { factor: 0.5 },
8690                },
8691                &ReferenceWeights::new(),
8692            )
8693            .unwrap(),
8694            (Some(vec![1.0, 1.0e30]), 1.0)
8695        );
8696
8697        // YaRN factors resolve to the transformers-twin divisors + attention factor
8698        // (values pinned in memra-gguf's yarn_divisors test against the banked receipt).
8699        let (yarn_factors, yarn_mscale) = rope_factor_values(
8700            &RopePlan {
8701                dimensions: 4,
8702                base: 10_000.0,
8703                factors: RopeFactors::Yarn {
8704                    factor: 2.0,
8705                    original_context: 8,
8706                    beta_fast: 32.0,
8707                    beta_slow: 1.0,
8708                },
8709            },
8710            &ReferenceWeights::new(),
8711        )
8712        .unwrap();
8713        let yarn_factors = yarn_factors.unwrap();
8714        assert_eq!(yarn_factors[0], 1.0);
8715        assert!((yarn_factors[1] - 2.0).abs() < 1e-6);
8716        assert!((yarn_mscale - 1.069_314_7).abs() < 1e-6);
8717    }
8718
8719    /// Every expected number below is hand-derived from the modular_qwen4_exp.py math
8720    /// (SEMANTICS.md §Gated residual), NOT read back from the code under test.
8721    #[test]
8722    // excessive_precision: the assert literals quote the hand derivation digits verbatim.
8723    #[allow(clippy::excessive_precision)]
8724    fn gated_residual_read_and_write_match_hand_derived_two_stream_toy() {
8725        let (streams, hidden, rank, tokens) = (2usize, 2usize, 1usize, 1usize);
8726        let wide = streams * hidden;
8727        let prefix = "trunk.layers.0.";
8728        let sublayer = "attn_hyper_connection.";
8729        let insert =
8730            |weights: &mut ReferenceWeights, suffix: &str, shape: &[usize], data: &[f32]| {
8731                weights.insert(
8732                    qwen4exp_family_id(format!("{prefix}{sublayer}{suffix}")),
8733                    weight(shape, data),
8734                );
8735            };
8736        // x = [3,4 | 6,8]: both stream groups are parallel, so grouped normalization maps
8737        // them to the SAME direction — n = (3,4)/sqrt(12.5+1e-6) per group. That equality
8738        // is itself the group-independence assertion.
8739        let x = [3.0, 4.0, 6.0, 8.0];
8740
8741        // Case A: zero down/up/inject weights => w = sigmoid(0) = 0.5 everywhere and
8742        // inject = 2*sigmoid(0) = 1; mixed[c] = 0.5*(n0[c]+n1[c])/2.
8743        let mut weights = ReferenceWeights::new();
8744        insert(&mut weights, "hc_norm.weight", &[wide], &[1.0; 4]);
8745        insert(
8746            &mut weights,
8747            "input_mix_weight_down.weight",
8748            &[rank, wide],
8749            &[0.0; 4],
8750        );
8751        insert(
8752            &mut weights,
8753            "input_mix_weight_up.weight",
8754            &[wide, rank],
8755            &[0.0; 4],
8756        );
8757        insert(
8758            &mut weights,
8759            "block_inject_weight.weight",
8760            &[streams, wide],
8761            &[0.0; 8],
8762        );
8763        let (mixed, inject) = gated_residual_read(
8764            &weights, prefix, sublayer, &x, tokens, streams, hidden, rank, 1e-6, true,
8765        )
8766        .unwrap();
8767        // Hand: n = (0.84852810, 1.13137080); mixed = (0.42426406, 0.56568541).
8768        assert!((mixed[0] - 0.424_264_06).abs() < 1e-5, "{mixed:?}");
8769        assert!((mixed[1] - 0.565_685_41).abs() < 1e-5, "{mixed:?}");
8770        assert!((inject[0] - 1.0).abs() < 1e-6 && (inject[1] - 1.0).abs() < 1e-6);
8771
8772        // Case B: down = [1,0,0,0], up = ones, inject row0 = [1,0,0,0], row1 = 0.
8773        // low  = silu(n[0]/2)           = silu(0.42426406)   = 0.25646896
8774        // w    = sigmoid(low)           = 0.56376809 for every dim
8775        // mixed[c] = w*(n0[c]+n1[c])/2  = (0.47837307, 0.63783076)
8776        // inject   = (2*sigmoid(n[0]/2), 2*sigmoid(0)) = (1.20900630, 1.0)
8777        insert(
8778            &mut weights,
8779            "input_mix_weight_down.weight",
8780            &[rank, wide],
8781            &[1.0, 0.0, 0.0, 0.0],
8782        );
8783        insert(
8784            &mut weights,
8785            "input_mix_weight_up.weight",
8786            &[wide, rank],
8787            &[1.0; 4],
8788        );
8789        insert(
8790            &mut weights,
8791            "block_inject_weight.weight",
8792            &[streams, wide],
8793            &[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
8794        );
8795        let (mixed, inject) = gated_residual_read(
8796            &weights, prefix, sublayer, &x, tokens, streams, hidden, rank, 1e-6, true,
8797        )
8798        .unwrap();
8799        assert!((mixed[0] - 0.478_373_07).abs() < 1e-5, "{mixed:?}");
8800        assert!((mixed[1] - 0.637_830_76).abs() < 1e-5, "{mixed:?}");
8801        assert!((inject[0] - 1.208_999_4).abs() < 1e-4, "{inject:?}");
8802        assert!((inject[1] - 1.0).abs() < 1e-6, "{inject:?}");
8803
8804        // Write: out = PRE-norm wide + block_out ⊗ inject, block_out = (1, -1)
8805        // => (3+1.209, 4-1.209, 6+1, 8-1).
8806        let mut wide_state = x.to_vec();
8807        gated_residual_write(
8808            &mut wide_state,
8809            &[1.0, -1.0],
8810            &inject,
8811            tokens,
8812            streams,
8813            hidden,
8814        );
8815        assert!((wide_state[0] - 4.209_006_3).abs() < 1e-4, "{wide_state:?}");
8816        assert!((wide_state[1] - 2.790_993_7).abs() < 1e-4, "{wide_state:?}");
8817        assert!((wide_state[2] - 7.0).abs() < 1e-6, "{wide_state:?}");
8818        assert!((wide_state[3] - 7.0).abs() < 1e-6, "{wide_state:?}");
8819    }
8820
8821    /// GDN with the qwen4_exp sigmoid z-gate — the ONE divergence from qwen3_5. Single
8822    /// token, identity-shaped projections, gate logit 2.0:
8823    ///   conv (k=1, w=1) => q=k=(silu(1),0), v=(silu(2),0); l2norm makes q~=k unit;
8824    ///   beta=sigmoid(0)=0.5, one step from zero state => mixed = (k.q)*v*beta/sqrt(2);
8825    ///   rms_norm => (1.41420992, 0); out = norm * act(2).
8826    /// Hand: sigmoid arm (1.24563196, 0); silu arm would be (2.49126392, 0).
8827    #[test]
8828    // excessive_precision: the assert literals quote the hand derivation digits verbatim.
8829    #[allow(clippy::excessive_precision)]
8830    fn gdn_sigmoid_gate_matches_hand_derived_single_token() {
8831        use memra_gguf::model_plan::GatedDeltaNetPlan;
8832
8833        let hidden = 2usize;
8834        let mut weights = ReferenceWeights::new();
8835        weights.insert(
8836            layer_id(0, LayerTensor::GdnQkv),
8837            weight(
8838                &[6, 2],
8839                &[1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 0.0, 1.0, 2.0, 0.0, 0.0, 2.0],
8840            ),
8841        );
8842        weights.insert(
8843            layer_id(0, LayerTensor::GdnGate),
8844            weight(&[2, 2], &[2.0, 0.0, 0.0, 1.0]),
8845        );
8846        weights.insert(
8847            layer_id(0, LayerTensor::GdnBeta),
8848            weight(&[1, 2], &[0.0, 0.0]),
8849        );
8850        weights.insert(
8851            layer_id(0, LayerTensor::GdnAlpha),
8852            weight(&[1, 2], &[0.0, 0.0]),
8853        );
8854        weights.insert(layer_id(0, LayerTensor::GdnA), weight(&[1], &[0.0]));
8855        weights.insert(layer_id(0, LayerTensor::GdnDtBias), weight(&[1], &[0.0]));
8856        weights.insert(layer_id(0, LayerTensor::GdnNorm), weight(&[2], &[1.0, 1.0]));
8857        weights.insert(
8858            layer_id(0, LayerTensor::GdnConv1d),
8859            weight(&[6, 1], &[1.0; 6]),
8860        );
8861        weights.insert(
8862            layer_id(0, LayerTensor::GdnOutput),
8863            weight(&[2, 2], &[1.0, 0.0, 0.0, 1.0]),
8864        );
8865        let plan = GatedDeltaNetPlan {
8866            key_heads: 1,
8867            value_heads: 1,
8868            key_head_dim: 2,
8869            value_head_dim: 2,
8870            conv_kernel: 1,
8871            gate_activation: GdnGateActivation::Sigmoid,
8872        };
8873        let (sigmoid_out, _) =
8874            gated_delta_net(0, &plan, 1e-6, &weights, &[1.0, 0.0], 1, hidden).unwrap();
8875        assert!(
8876            (sigmoid_out[0] - 1.245_632_0).abs() < 1e-4,
8877            "{sigmoid_out:?}"
8878        );
8879        assert!(sigmoid_out[1].abs() < 1e-6, "{sigmoid_out:?}");
8880
8881        let silu_plan = GatedDeltaNetPlan {
8882            gate_activation: GdnGateActivation::Silu,
8883            ..plan
8884        };
8885        let (silu_out, _) =
8886            gated_delta_net(0, &silu_plan, 1e-6, &weights, &[1.0, 0.0], 1, hidden).unwrap();
8887        assert!((silu_out[0] - 2.491_263_9).abs() < 1e-4, "{silu_out:?}");
8888    }
8889
8890    /// Crafted 12-token sequence with an unambiguous top-k choice: only tokens 4..8 carry
8891    /// key mass (block 1); every other block pools to the ZERO vector, which rope and
8892    /// normalization preserve, so its relu score is exactly 0 while block 1 scores
8893    /// strictly positive (2*cos(Δpos) with Δpos ∈ {5, 7} rad, both cos > 0). budget = 1
8894    /// block. Also pins the tail rule, including the boundary case where a query's own
8895    /// unselected complete block makes the query NOT see itself.
8896    #[test]
8897    fn micro_block_indexer_selects_unambiguous_block_and_always_keeps_the_tail() {
8898        let tokens = 12usize;
8899        let hidden = 2usize;
8900        let overlay = MicroBlockIndexPlan {
8901            query_heads: 1,
8902            kv_heads: 1,
8903            head_dim: 2,
8904            rope_dimensions: 2,
8905            block_size: 4,
8906            budget_blocks: 1,
8907            budget_tokens: 4,
8908        };
8909        let rope = RopePlan {
8910            dimensions: 2,
8911            base: 10_000.0,
8912            factors: memra_gguf::model_plan::RopeFactors::None,
8913        };
8914        let prefix = "trunk.layers.0.";
8915        let mut weights = ReferenceWeights::new();
8916        // q rows = identity (q = x); k rows read only x[1] scaled by 10.
8917        weights.insert(
8918            qwen4exp_family_id(format!("{prefix}self_attn.indexer.index_qk_proj.weight")),
8919            weight(&[4, 2], &[1.0, 0.0, 0.0, 1.0, 0.0, 10.0, 0.0, 0.0]),
8920        );
8921        for norm in ["q_layernorm", "k_layernorm"] {
8922            weights.insert(
8923                qwen4exp_family_id(format!("{prefix}self_attn.indexer.{norm}.weight")),
8924                weight(&[2], &[1.0, 1.0]),
8925            );
8926        }
8927        let mut x = vec![0.0; tokens * hidden];
8928        for token in 0..tokens {
8929            x[token * hidden] = 1.0; // every query is (1, 0)
8930            if (4..8).contains(&token) {
8931                x[token * hidden + 1] = 1.0; // block-1 keys become (10, 0)
8932            }
8933        }
8934        let mask = micro_block_selection_mask(
8935            0, &overlay, &rope, 1e-6, &weights, prefix, &x, tokens, hidden,
8936        )
8937        .unwrap();
8938        let row = |token: usize| &mask[token * tokens..(token + 1) * tokens];
8939        // t=0: no complete block, tail = {0}.
8940        assert_eq!(
8941            row(0),
8942            &[
8943                true, false, false, false, false, false, false, false, false, false, false, false
8944            ]
8945        );
8946        // t=5: one complete block (0..4, the only candidate) + tail {4,5}.
8947        assert_eq!(
8948            row(5),
8949            &[
8950                true, true, true, true, true, true, false, false, false, false, false, false
8951            ]
8952        );
8953        // t=9: blocks {0,1} complete, block 1 wins (score>0 vs 0), tail {8,9}.
8954        assert_eq!(
8955            row(9),
8956            &[
8957                false, false, false, false, true, true, true, true, true, true, false, false
8958            ]
8959        );
8960        // t=11: blocks {0,1,2} complete, tail EMPTY; only block 1 selected — the query
8961        // does not even see itself (blocks-only selection at exact boundaries).
8962        assert_eq!(
8963            row(11),
8964            &[
8965                false, false, false, false, true, true, true, true, false, false, false, false
8966            ]
8967        );
8968    }
8969
8970    /// The selection mask gates full attention to exactly the selected sources: a query
8971    /// restricted to itself must return its own VALUE row bit-for-bit reasoning
8972    /// (softmax over one score = 1), with identity projections output == input.
8973    #[test]
8974    fn full_attention_selection_mask_restricts_sources_to_hand_derived_rows() {
8975        use memra_gguf::model_plan::{FullAttentionPlan, RopeFactors, TensorPresence};
8976
8977        let plan = FullAttentionPlan {
8978            query_heads: 1,
8979            kv_heads: 1,
8980            key_head_dim: 2,
8981            value_head_dim: 2,
8982            rope: RopePlan {
8983                dimensions: 2,
8984                base: 10_000.0,
8985                factors: RopeFactors::None,
8986            },
8987            qk_norm: TensorPresence::Absent,
8988            output_gate: memra_gguf::config::AttentionGateKind::None,
8989            scale: AttentionScale::InverseSqrtKeyDim,
8990            value_projection: ValueProjection::Separate,
8991            value_norm: ValueNorm::None,
8992        };
8993        let identity = [1.0, 0.0, 0.0, 1.0];
8994        let mut weights = ReferenceWeights::new();
8995        for tensor in [
8996            LayerTensor::Query,
8997            LayerTensor::Key,
8998            LayerTensor::Value,
8999            LayerTensor::AttentionOutput,
9000        ] {
9001            weights.insert(layer_id(0, tensor), weight(&[2, 2], &identity));
9002        }
9003        // Orthogonal rows keep the causal softmax far from saturation (score gap ~1.3),
9004        // so the unmasked row visibly mixes ~21% of source 0.
9005        let x = [1.0, 0.0, 0.0, 1.0];
9006        let diagonal = [true, false, false, true];
9007        let (masked, _) =
9008            full_attention(0, &plan, None, 1e-6, &weights, &x, 2, 2, Some(&diagonal)).unwrap();
9009        for index in 0..4 {
9010            assert!((masked[index] - x[index]).abs() < 1e-6, "{masked:?}");
9011        }
9012        let (unmasked, _) = full_attention(0, &plan, None, 1e-6, &weights, &x, 2, 2, None).unwrap();
9013        assert!(
9014            (unmasked[2] - x[2]).abs() > 1e-3,
9015            "causal row must mix sources"
9016        );
9017
9018        let starving = [true, false, false, false];
9019        let error =
9020            full_attention(0, &plan, None, 1e-6, &weights, &x, 2, 2, Some(&starving)).unwrap_err();
9021        assert!(matches!(error, ReferenceError::InvalidPlan { .. }));
9022    }
9023
9024    /// N-gram id math recomputed independently below (wrapping i64 multiply, XOR, floor
9025    /// mod, offset — SEMANTICS.md §PLE), with a multiplier big enough that the product
9026    /// wraps negative and exercises the floor-mod arm.
9027    #[test]
9028    fn ngram_ids_match_independently_computed_hash_chain() {
9029        let multipliers = [0x4000_0000_0000_0001_i64, 1_000_003, 7_777_777];
9030        let sizes = [97_i64, 89, 83, 79];
9031        let offsets = [0_i64, 97, 186, 269];
9032        let (max_ngram, heads_per_ngram, eos) = (3usize, 2usize, 9u32);
9033        let token_ids = [5u32, 7];
9034        let ids = ngram_ids(
9035            &token_ids,
9036            &multipliers,
9037            &sizes,
9038            &offsets,
9039            max_ngram,
9040            heads_per_ngram,
9041            eos,
9042            0,
9043        )
9044        .unwrap();
9045
9046        // history = [9, 9, 5, 7]; shifted[1] = [9,9,9,5]; shifted[2] = [9,9,9,9]
9047        // (the two context positions read EOS; position 3 shifted-by-1 reads token 5).
9048        let expect = |mixed: i64, head: usize| mixed.rem_euclid(sizes[head]) + offsets[head];
9049        let bigram_t0 = 5_i64.wrapping_mul(multipliers[0]) ^ 9_i64.wrapping_mul(multipliers[1]);
9050        let trigram_t0 = bigram_t0 ^ 9_i64.wrapping_mul(multipliers[2]);
9051        let bigram_t1 = 7_i64.wrapping_mul(multipliers[0]) ^ 5_i64.wrapping_mul(multipliers[1]);
9052        let trigram_t1 = bigram_t1 ^ 9_i64.wrapping_mul(multipliers[2]);
9053        // 7 * (2^62 + 1) wraps to 2^63 + 2^62 + 7, i.e. negative i64; floor mod must
9054        // still land non-negative (torch.remainder semantics).
9055        assert!(7_i64.wrapping_mul(multipliers[0]) < 0);
9056        assert_eq!(
9057            ids,
9058            vec![
9059                expect(bigram_t0, 0),
9060                expect(bigram_t0, 1),
9061                expect(trigram_t0, 2),
9062                expect(trigram_t0, 3),
9063                expect(bigram_t1, 0),
9064                expect(bigram_t1, 1),
9065                expect(trigram_t1, 2),
9066                expect(trigram_t1, 3),
9067            ]
9068        );
9069        assert!(ids.iter().all(|&id| id >= 0));
9070    }
9071
9072    /// Hand-derived shift vectors for history [E,E,5,6,E,7,8] (E = 63):
9073    ///   eos strictly-before: [-1,0,1,1,1,4,4]; segment starts [0,1,2,2,2,5,5];
9074    ///   in-segment positions [0,0,0,1,2,0,1].
9075    /// shift=1 keeps positions {3,4,6} (note position 4 — the EOS itself — reads 6, its
9076    /// in-segment index counts within the PREVIOUS segment); shift=2 keeps only {4}.
9077    #[test]
9078    fn eos_segment_reset_reads_eos_across_boundaries() {
9079        let eos = 63i64;
9080        let history = [eos, eos, 5, 6, eos, 7, 8];
9081        assert_eq!(shift_right_ignore_eos(&history, 0, eos), history.to_vec());
9082        assert_eq!(
9083            shift_right_ignore_eos(&history, 1, eos),
9084            vec![eos, eos, eos, 5, 6, eos, 7]
9085        );
9086        assert_eq!(
9087            shift_right_ignore_eos(&history, 2, eos),
9088            vec![eos, eos, eos, eos, 5, eos, eos]
9089        );
9090    }
9091
9092    /// Scalar-channel PLE block pinning the gather -> gate -> dilated-conv chain by hand:
9093    /// wide stream 0 => query norm 0 => gate = sigmoid(0) = 0.5, so gated = 0.5*value;
9094    /// normed scalars n_t = g_t/sqrt(g_t^2+1e-6); conv (kernel 2, dilation = max_ngram
9095    /// = 2, taps w = [10, 1]) reads out[t] = g_t + silu(10*n_{t-2} + n_t) with the
9096    /// out-of-range tap dropped. Hand values below; a REVERSED tap order would give
9097    /// out[2] = 11.5068593 instead of 8.8685460, so this pins conv orientation AND the
9098    /// dilation reach (t-2, not t-1).
9099    #[test]
9100    // excessive_precision: the assert literals quote the hand derivation digits verbatim.
9101    #[allow(clippy::excessive_precision)]
9102    fn ple_block_matches_hand_derived_scalar_gather_gate_and_dilated_conv() {
9103        let prefix = "trunk.layers.1.";
9104        let mut weights = ReferenceWeights::new();
9105        let family = |suffix: &str| qwen4exp_family_id(format!("{prefix}{suffix}"));
9106        weights.insert(
9107            family("ple.ple_embedding.layer_multipliers"),
9108            ReferenceTensor::new_i64(vec![2], vec![1, 0]).unwrap(),
9109        );
9110        weights.insert(
9111            family("ple.ple_embedding.ngram_heads_vocab_sizes"),
9112            ReferenceTensor::new_i64(vec![1], vec![5]).unwrap(),
9113        );
9114        weights.insert(
9115            family("ple.ple_embedding.ngram_heads_offsets"),
9116            ReferenceTensor::new_i64(vec![1], vec![0]).unwrap(),
9117        );
9118        // ids = token mod 5 = [1, 2, 3] -> values [0.002, 0.4, 1.6]
9119        weights.insert(
9120            family("ple.ple_embedding.ngram_embedding"),
9121            weight(&[5, 1], &[0.0, 0.002, 0.4, 1.6, 0.0]),
9122        );
9123        weights.insert(family("ple.key_proj.weight"), weight(&[1, 1], &[1.0]));
9124        weights.insert(family("ple.value_proj.weight"), weight(&[1, 1], &[1.0]));
9125        for norm in ["norm_key", "norm_query", "norm_conv"] {
9126            weights.insert(family(&format!("ple.{norm}.weight")), weight(&[1], &[1.0]));
9127        }
9128        weights.insert(family("ple.conv1d.weight"), weight(&[1, 2], &[10.0, 1.0]));
9129        let plan = memra_gguf::model_plan::PleEmbeddingPlan {
9130            ngram_heads: 1,
9131            head_embed_dim: 1,
9132            vocab_shards: 1,
9133            embed_dim: 1,
9134            conv_kernel: 2,
9135            max_ngram: 2,
9136            eos_token_id: 4,
9137        };
9138        let wide_state = [0.0; 3];
9139        let output = ple_block(
9140            1,
9141            &plan,
9142            1e-6,
9143            &weights,
9144            prefix,
9145            &wide_state,
9146            &[1, 2, 3],
9147            3,
9148            1,
9149            1,
9150        )
9151        .unwrap();
9152        assert!((output[0] - 0.474_592_9).abs() < 1e-4, "{output:?}");
9153        assert!((output[1] - 0.931_047_0).abs() < 1e-4, "{output:?}");
9154        assert!((output[2] - 8.868_546_0).abs() < 1e-3, "{output:?}");
9155    }
9156
9157    /// The qwen4_exp pack's tiny plan executes end-to-end through `execute`: gated
9158    /// residual entry/exit, GDN + QSA trunk, PLE on layer 1, MoE with the gated shared
9159    /// expert, and the separate-projection MTP block. 16 tokens so the indexer budget
9160    /// (2 blocks) actually BINDS (3-4 complete blocks at the last queries).
9161    #[test]
9162    fn qwen4exp_tiny_plan_executes_gated_residual_qsa_ple_moe_and_mtp() {
9163        let pack = memra_gguf::model_packs::by_alias("qwen4_exp").expect("qwen4_exp pack");
9164        let plan = pack.compile_tiny_plan().expect("tiny plan compiles");
9165        assert_eq!(plan.layers.len(), 4);
9166        assert_eq!(plan.mtp_blocks.len(), 1);
9167        let fixture = deterministic_fixture(&plan).unwrap();
9168        assert!(
9169            !fixture.weights.contains_key(&TensorId::OutputNorm),
9170            "exit-mixer plans must not fabricate a final norm"
9171        );
9172        let token_ids: Vec<u32> = (1..=16).collect();
9173        let first = execute(&plan, &fixture.weights, &token_ids).unwrap();
9174        let second = execute(&plan, &fixture.weights, &token_ids).unwrap();
9175        assert_eq!(first, second, "reference must be bit-deterministic");
9176        assert_eq!((first.tokens, first.vocab), (16, 64));
9177        assert!(first.logits.iter().all(|value| value.is_finite()));
9178        for (index, state) in first.state.layers.iter().enumerate() {
9179            if index == 3 {
9180                assert!(matches!(state, ReferenceLayerState::Kv { .. }));
9181            } else {
9182                assert!(matches!(state, ReferenceLayerState::Recurrent { .. }));
9183            }
9184        }
9185        // MTP: wide (streams*hidden) post-layer state is the K>1 carrier.
9186        assert_eq!(first.mtp.len(), 1);
9187        assert_eq!(first.mtp[0].hidden.len(), 16 * 2 * 16);
9188        assert_eq!(first.mtp[0].logits.len(), 16 * 64);
9189        assert!(first.mtp[0].logits.iter().all(|value| value.is_finite()));
9190
9191        // The QSA selection binds: zeroing the indexer projection makes every block
9192        // score exactly 0, so the pinned tie rule keeps the LOWEST-indexed blocks —
9193        // a different selection than the trained-shaped fixture picks. (A sign flip
9194        // would NOT work here: negating q and k together preserves every score.)
9195        let mut perturbed = fixture.weights.clone();
9196        perturbed
9197            .get_mut(&qwen4exp_family_id(
9198                "trunk.layers.3.self_attn.indexer.index_qk_proj.weight".into(),
9199            ))
9200            .expect("trunk indexer weights")
9201            .data
9202            .fill(0.0);
9203        let reindexed = execute(&plan, &perturbed, &token_ids).unwrap();
9204        assert_ne!(
9205            first.logits, reindexed.logits,
9206            "indexer selection must gate attention"
9207        );
9208
9209        // PLE binds: a different n-gram table moves the logits.
9210        let mut retabled = fixture.weights.clone();
9211        retabled
9212            .get_mut(&qwen4exp_family_id(
9213                "trunk.layers.1.ple.ple_embedding.ngram_embedding".into(),
9214            ))
9215            .expect("ngram table")
9216            .data
9217            .fill(0.25);
9218        let regathered = execute(&plan, &retabled, &token_ids).unwrap();
9219        assert_ne!(
9220            first.logits, regathered.logits,
9221            "PLE gather must feed layer 1"
9222        );
9223
9224        // The sigmoid-gated shared expert binds (MoE deliverable check).
9225        let mut regated = fixture.weights.clone();
9226        regated
9227            .get_mut(&layer_id(0, LayerTensor::SharedMlpInputGate))
9228            .expect("shared expert gate")
9229            .data
9230            .fill(4.0);
9231        let reshared = execute(&plan, &regated, &token_ids).unwrap();
9232        assert_ne!(
9233            first.logits, reshared.logits,
9234            "shared-expert sigmoid gate must scale the shared branch"
9235        );
9236    }
9237
9238    #[test]
9239    fn dense_gemma_executes_scaled_parallel_residual_and_k_as_v() {
9240        let config = ModelConfig::from_hf(&HfConfig::parse(
9241            r#"{"model_type":"gemma4","num_hidden_layers":2,"hidden_size":8,
9242            "num_attention_heads":2,"num_key_value_heads":1,
9243            "num_global_key_value_heads":1,"head_dim":4,"global_head_dim":4,
9244            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
9245            "rms_norm_eps":0.000001,"sliding_window":2,
9246            "final_logit_softcapping":30,
9247            "layer_types":["sliding_attention","full_attention"],
9248            "rope_parameters":{"full_attention":{"rope_theta":1000000,
9249            "partial_rotary_factor":0.5},"sliding_attention":{"rope_theta":10000}}}"#,
9250        ));
9251        let plan = ModelPlan::compile(&config).unwrap();
9252        assert_eq!(plan.embedding_scale, 8.0f32.sqrt());
9253        let fixture = deterministic_fixture(&plan).unwrap();
9254        assert!(
9255            !fixture
9256                .weights
9257                .contains_key(&layer_id(1, LayerTensor::Value))
9258        );
9259        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
9260        assert!(output.logits.iter().all(|value| value.is_finite()));
9261        let ReferenceLayerState::Kv { window, .. } = output.state.layers[0] else {
9262            panic!("expected SWA state");
9263        };
9264        assert_eq!(window, Some(2));
9265        let ReferenceLayerState::Kv { window, .. } = output.state.layers[1] else {
9266            panic!("expected global state");
9267        };
9268        assert_eq!(window, None);
9269        assert_eq!(
9270            output.logits[..4]
9271                .iter()
9272                .map(|value| value.to_bits())
9273                .collect::<Vec<_>>(),
9274            vec![3_198_203_366, 1_057_194_687, 3_185_247_713, 3_204_119_266]
9275        );
9276    }
9277
9278    /// ONE TensorId, ONE byte order. `deterministic_fixture` mints the reference's weights and
9279    /// `TensorContract` names the same ids for the engine's loader; any engine-vs-reference gate
9280    /// serves ONE set of bytes under both. A shape disagreement preserves element counts, so
9281    /// nothing else catches it — `MlaKeyUp` was minted `[head][nope][rank]` against the
9282    /// contract's `[head][rank][nope]` and every MLA parity comparison silently mis-strided the
9283    /// absorb operand (glm53-flash lane, 2026-08-28: 1.24e-1 relative on a micro fixture that
9284    /// drops to 6.9e-7 once the layouts agree).
9285    #[test]
9286    fn the_mla_fixture_shapes_match_the_tensor_contract() {
9287        use memra_gguf::tensor_contract::{
9288            CheckpointDialect, ContractOptions, OutputHead, TensorContract,
9289        };
9290
9291        use memra_gguf::model_plan::{MlaAttentionPlan, StatePlan};
9292
9293        // EVERY MLA extent DISTINCT (heads 2, q_lora 3, kv_rank 6, nope 4, v 5). The shared tiny
9294        // plan runs 2/4/4/4/4, where a transposed key plane has the SAME shape as a correct one
9295        // and this pin would be a tautology.
9296        let mut plan = kpool_mla_reference_plan();
9297        let AttentionPlan::Mla(MlaAttentionPlan::LatentKv {
9298            q_lora_rank,
9299            kv_lora_rank,
9300            qk_head_dim,
9301            value_head_dim,
9302            ..
9303        }) = &mut plan.layers[1].attention
9304        else {
9305            panic!("layer 1 of the tiny plan must be MLA LatentKv");
9306        };
9307        *q_lora_rank = 3;
9308        *kv_lora_rank = 6;
9309        *qk_head_dim = 4;
9310        *value_head_dim = 5;
9311        plan.layers[1].state = StatePlan::LatentKvCache {
9312            width: 6,
9313            index_width: 8,
9314        };
9315        let fixture = deterministic_fixture(&plan).unwrap();
9316        let contract = TensorContract::for_plan(
9317            &plan,
9318            CheckpointDialect::Gguf,
9319            ContractOptions {
9320                output_head: OutputHead::TiedToEmbedding,
9321            },
9322        )
9323        .unwrap();
9324        let mut checked = 0;
9325        for requirement in &contract.requirements {
9326            let Some(tensor) = fixture.weights.get(&requirement.id) else {
9327                continue;
9328            };
9329            // GGUF `ne` is fastest-axis-first; the fixture states row-major shapes.
9330            let mut wanted: Vec<usize> = requirement.shape.iter().map(|&d| d as usize).collect();
9331            wanted.reverse();
9332            let TensorId::Layer { tensor: kind, .. } = requirement.id else {
9333                continue;
9334            };
9335            if !matches!(
9336                kind,
9337                LayerTensor::MlaKeyUp | LayerTensor::MlaValueUp | LayerTensor::MlaQueryUp
9338            ) {
9339                continue;
9340            }
9341            assert_eq!(
9342                tensor.shape, wanted,
9343                "{:?}: fixture shape {:?} but the contract declares ne {:?}",
9344                requirement.id, tensor.shape, requirement.shape
9345            );
9346            checked += 1;
9347        }
9348        assert!(checked >= 3, "the plan must exercise the MLA planes");
9349    }
9350
9351    /// The glm5_next-shaped tiny plan: one KDA layer, one MLA+k-pool-indexer layer, sigmoid MoE,
9352    /// hyper-connections with the mean collapse. Shared by the execution gate and the
9353    /// fixture-vs-contract shape pin so both describe the SAME plan.
9354    fn kpool_mla_reference_plan() -> ModelPlan {
9355        use memra_gguf::model_plan::{
9356            DenseMlpPlan, KimiDeltaNetPlan, KpoolPlan, MlaAttentionPlan, MoeMlpPlan, RopeFactors,
9357            RopePlan, RouterPlan, SharedMlpPlan, SparseIndexPlan, StatePlan,
9358        };
9359
9360        let config = ModelConfig::from_hf(&HfConfig::parse(
9361            r#"{"model_type":"qwen3","num_hidden_layers":2,"hidden_size":8,
9362            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
9363            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
9364            "rms_norm_eps":0.00001}"#,
9365        ));
9366        let mut plan = ModelPlan::compile(&config).unwrap();
9367        plan.layers[0].attention = AttentionPlan::KimiDeltaNet(KimiDeltaNetPlan {
9368            num_heads: 2,
9369            head_dim: 4,
9370            conv_kernel: 3,
9371            gate_lower_bound: -5.0,
9372        });
9373        plan.layers[0].state = StatePlan::Recurrent {
9374            conv_width: 24,
9375            conv_kernel: 3,
9376            state_width: 32,
9377        };
9378        plan.layers[0].mlp = MlpPlan::Dense(DenseMlpPlan {
9379            intermediate_size: 16,
9380            activation: ActivationPlan::SwiGluPreClamped { limit: 10.0 },
9381        });
9382        plan.layers[1].attention = AttentionPlan::Mla(MlaAttentionPlan::LatentKv {
9383            query_heads: 2,
9384            q_lora_rank: 4,
9385            kv_lora_rank: 4,
9386            qk_head_dim: 4,
9387            rope_head_dim: 0,
9388            value_head_dim: 4,
9389            rope: RopePlan {
9390                dimensions: 0,
9391                base: 10_000.0,
9392                factors: RopeFactors::None,
9393            },
9394            sparse_index: SparseIndexPlan::Own {
9395                heads: 2,
9396                head_dim: 4,
9397                top_k: 4,
9398                kpool: Some(KpoolPlan {
9399                    pool: 2,
9400                    always_select_tail: true,
9401                }),
9402            },
9403        });
9404        plan.layers[1].state = StatePlan::LatentKvCache {
9405            width: 4,
9406            index_width: 8,
9407        };
9408        plan.layers[1].mlp = MlpPlan::Moe(MoeMlpPlan {
9409            expert_count: 4,
9410            experts_per_token: 2,
9411            expert_intermediate_size: 4,
9412            router: RouterPlan::Sigmoid {
9413                normalize_selected: true,
9414                scaling_factor: 2.5,
9415                selection_bias: true,
9416            },
9417            shared: Some(SharedMlpPlan {
9418                intermediate_size: 4,
9419                gated: false,
9420            }),
9421            activation: ActivationPlan::SwiGluPreClamped { limit: 10.0 },
9422        });
9423        for layer in &mut plan.layers {
9424            layer.residual = ResidualTopology::HyperConnections {
9425                streams: 2,
9426                epsilon: 1e-6,
9427                sinkhorn_iterations: 2,
9428                collapse: HcCollapse::Mean,
9429            };
9430        }
9431        plan
9432    }
9433
9434    #[test]
9435    fn glm5_shaped_tiny_plan_executes_kda_kpool_mla_and_mean_collapse_deterministically() {
9436        let plan = kpool_mla_reference_plan();
9437        let fixture = deterministic_fixture(&plan).unwrap();
9438        // The mean collapse owns no learned head tensors.
9439        assert!(!fixture.weights.contains_key(&TensorId::HyperHeadFunction));
9440        assert!(
9441            fixture
9442                .weights
9443                .contains_key(&layer_id(1, LayerTensor::SparseCompressorGate))
9444        );
9445        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
9446        assert_eq!(output.logits.len(), fixture.token_ids.len() * 32);
9447        assert!(output.logits.iter().all(|value| value.is_finite()));
9448        assert!(matches!(
9449            output.state.layers[0],
9450            ReferenceLayerState::Recurrent { conv_width: 24, .. }
9451        ));
9452        assert!(matches!(
9453            output.state.layers[1],
9454            ReferenceLayerState::LatentKv { width: 4, .. }
9455        ));
9456        let second = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
9457        assert_eq!(
9458            output
9459                .logits
9460                .iter()
9461                .map(|value| value.to_bits())
9462                .collect::<Vec<_>>(),
9463            second
9464                .logits
9465                .iter()
9466                .map(|value| value.to_bits())
9467                .collect::<Vec<_>>()
9468        );
9469    }
9470
9471    #[test]
9472    fn kimi_delta_net_matches_hand_derived_three_token_recurrence() {
9473        use memra_gguf::model_plan::KimiDeltaNetPlan;
9474
9475        let plan = KimiDeltaNetPlan {
9476            num_heads: 1,
9477            head_dim: 2,
9478            conv_kernel: 2,
9479            gate_lower_bound: -5.0,
9480        };
9481        let x = [[0.5f32, -0.3], [0.1, 0.8], [-0.6, 0.2]];
9482        let wq = [[0.7f32, -0.2], [0.3, 0.5]];
9483        let wk = [[0.4f32, 0.1], [-0.3, 0.6]];
9484        let wv = [[0.9f32, 0.2], [-0.1, 0.8]];
9485        let q_conv = [[0.3f32, 0.7], [-0.2, 0.9]];
9486        let k_conv = [[0.5f32, 0.5], [0.1, 0.8]];
9487        let v_conv = [[0.2f32, 0.6], [0.4, 0.4]];
9488        let f_a = [[0.6f32, -0.4], [0.2, 0.3]];
9489        let f_b = [[0.5f32, 0.1], [-0.2, 0.7]];
9490        let dt_bias = [0.05f32, -0.1];
9491        let a_log = [0.2f32];
9492        let b_proj = [[0.4f32, -0.6]];
9493        let g_a = [[0.3f32, 0.2], [-0.5, 0.4]];
9494        let g_b = [[0.6f32, -0.3], [0.2, 0.5]];
9495        let o_norm = [1.0f32, 1.5];
9496        let wo = [[0.8f32, -0.4], [0.3, 0.9]];
9497
9498        let mut weights = ReferenceWeights::new();
9499        let flat = |rows: &[[f32; 2]]| -> Vec<f32> { rows.iter().flatten().copied().collect() };
9500        weights.insert(
9501            layer_id(0, LayerTensor::KdaQuery),
9502            weight(&[2, 2], &flat(&wq)),
9503        );
9504        weights.insert(
9505            layer_id(0, LayerTensor::KdaKey),
9506            weight(&[2, 2], &flat(&wk)),
9507        );
9508        weights.insert(
9509            layer_id(0, LayerTensor::KdaValue),
9510            weight(&[2, 2], &flat(&wv)),
9511        );
9512        weights.insert(
9513            layer_id(0, LayerTensor::KdaQueryConv),
9514            weight(&[2, 2], &flat(&q_conv)),
9515        );
9516        weights.insert(
9517            layer_id(0, LayerTensor::KdaKeyConv),
9518            weight(&[2, 2], &flat(&k_conv)),
9519        );
9520        weights.insert(
9521            layer_id(0, LayerTensor::KdaValueConv),
9522            weight(&[2, 2], &flat(&v_conv)),
9523        );
9524        weights.insert(
9525            layer_id(0, LayerTensor::KdaForgetDown),
9526            weight(&[2, 2], &flat(&f_a)),
9527        );
9528        weights.insert(
9529            layer_id(0, LayerTensor::KdaForgetUp),
9530            weight(&[2, 2], &flat(&f_b)),
9531        );
9532        weights.insert(layer_id(0, LayerTensor::KdaDtBias), weight(&[2], &dt_bias));
9533        weights.insert(layer_id(0, LayerTensor::KdaALog), weight(&[1], &a_log));
9534        weights.insert(
9535            layer_id(0, LayerTensor::KdaBeta),
9536            weight(&[1, 2], &flat(&b_proj)),
9537        );
9538        weights.insert(
9539            layer_id(0, LayerTensor::KdaGateDown),
9540            weight(&[2, 2], &flat(&g_a)),
9541        );
9542        weights.insert(
9543            layer_id(0, LayerTensor::KdaGateUp),
9544            weight(&[2, 2], &flat(&g_b)),
9545        );
9546        weights.insert(
9547            layer_id(0, LayerTensor::KdaOutputNorm),
9548            weight(&[2], &o_norm),
9549        );
9550        weights.insert(
9551            layer_id(0, LayerTensor::KdaOutput),
9552            weight(&[2, 2], &flat(&wo)),
9553        );
9554
9555        let x_flat: Vec<f32> = x.iter().flatten().copied().collect();
9556        let (output, _) = kimi_delta_net(0, &plan, 1e-5, &weights, &x_flat, 3, 2).unwrap();
9557
9558        // Duplicated arithmetic, written independently of the operator.
9559        let sig = |value: f32| 1.0 / (1.0 + (-value).exp());
9560        let act = |value: f32| value * (1.0 / (1.0 + (-value).exp()));
9561        let mat2 = |m: &[[f32; 2]; 2], v: [f32; 2]| {
9562            [
9563                m[0][0] * v[0] + m[0][1] * v[1],
9564                m[1][0] * v[0] + m[1][1] * v[1],
9565            ]
9566        };
9567        let mut q_proj = [[0.0f32; 2]; 3];
9568        let mut k_proj = [[0.0f32; 2]; 3];
9569        let mut v_proj = [[0.0f32; 2]; 3];
9570        for token in 0..3 {
9571            q_proj[token] = mat2(&wq, x[token]);
9572            k_proj[token] = mat2(&wk, x[token]);
9573            v_proj[token] = mat2(&wv, x[token]);
9574        }
9575        let causal_conv = |proj: &[[f32; 2]; 3], conv: &[[f32; 2]; 2]| {
9576            let mut out = [[0.0f32; 2]; 3];
9577            for token in 0..3 {
9578                for channel in 0..2 {
9579                    let previous = if token == 0 {
9580                        0.0
9581                    } else {
9582                        proj[token - 1][channel]
9583                    };
9584                    out[token][channel] =
9585                        act(conv[channel][0] * previous + conv[channel][1] * proj[token][channel]);
9586                }
9587            }
9588            out
9589        };
9590        let mut q = causal_conv(&q_proj, &q_conv);
9591        let mut k = causal_conv(&k_proj, &k_conv);
9592        let v = causal_conv(&v_proj, &v_conv);
9593        for token in 0..3 {
9594            let q_inv = 1.0 / (q[token][0] * q[token][0] + q[token][1] * q[token][1] + 1e-6).sqrt();
9595            let k_inv = 1.0 / (k[token][0] * k[token][0] + k[token][1] * k[token][1] + 1e-6).sqrt();
9596            for channel in 0..2 {
9597                q[token][channel] *= q_inv * (1.0 / 2.0f32.sqrt());
9598                k[token][channel] *= k_inv;
9599            }
9600        }
9601        let decay_rate = a_log[0].exp();
9602        let mut expected = Vec::new();
9603        let mut state = [[0.0f32; 2]; 2];
9604        for token in 0..3 {
9605            let f_lin = mat2(&f_b, mat2(&f_a, x[token]));
9606            let g = [
9607                -5.0 * sig(decay_rate * (f_lin[0] + dt_bias[0])),
9608                -5.0 * sig(decay_rate * (f_lin[1] + dt_bias[1])),
9609            ];
9610            let beta = sig(b_proj[0][0] * x[token][0] + b_proj[0][1] * x[token][1]);
9611            for key_index in 0..2 {
9612                #[allow(clippy::needless_range_loop)]
9613                // allow: the explicit index loop keeps the offset arithmetic visible and aligned with the device-side indexing
9614                for value_index in 0..2 {
9615                    state[key_index][value_index] *= g[key_index].exp();
9616                }
9617            }
9618            let mut core = [0.0f32; 2];
9619            for value_index in 0..2 {
9620                let memory =
9621                    state[0][value_index] * k[token][0] + state[1][value_index] * k[token][1];
9622                let delta = (v[token][value_index] - memory) * beta;
9623                state[0][value_index] += k[token][0] * delta;
9624                state[1][value_index] += k[token][1] * delta;
9625            }
9626            for value_index in 0..2 {
9627                core[value_index] =
9628                    state[0][value_index] * q[token][0] + state[1][value_index] * q[token][1];
9629            }
9630            let gate = mat2(&g_b, mat2(&g_a, x[token]));
9631            let mean_square = (core[0] * core[0] + core[1] * core[1]) / 2.0;
9632            let inverse = 1.0 / (mean_square + 1e-5).sqrt();
9633            let gated = [
9634                core[0] * inverse * o_norm[0] * sig(gate[0]),
9635                core[1] * inverse * o_norm[1] * sig(gate[1]),
9636            ];
9637            let final_row = mat2(&wo, gated);
9638            expected.extend_from_slice(&final_row);
9639        }
9640        assert_eq!(output.len(), expected.len());
9641        for (index, (actual, wanted)) in output.iter().zip(&expected).enumerate() {
9642            assert!(
9643                (actual - wanted).abs() < 1e-5,
9644                "output[{index}] = {actual}, expected {wanted}"
9645            );
9646        }
9647    }
9648
9649    #[test]
9650    fn kpool_indexer_selects_causal_pools_and_appends_visible_tail() {
9651        use memra_gguf::model_plan::KpoolPlan;
9652
9653        let tokens = 8;
9654        let hidden = 2;
9655        let q_rank = 2;
9656        let identity = [1.0f32, 0.0, 0.0, 1.0];
9657        let mut weights = ReferenceWeights::new();
9658        weights.insert(
9659            layer_id(0, LayerTensor::SparseQuery),
9660            weight(&[2, 2], &identity),
9661        );
9662        weights.insert(
9663            layer_id(0, LayerTensor::SparseKey),
9664            weight(&[2, 2], &identity),
9665        );
9666        weights.insert(
9667            layer_id(0, LayerTensor::SparseKeyNorm),
9668            weight(&[2], &[1.0, 1.0]),
9669        );
9670        weights.insert(
9671            layer_id(0, LayerTensor::SparseKeyNormBias),
9672            weight(&[2], &[0.0, 0.0]),
9673        );
9674        weights.insert(
9675            layer_id(0, LayerTensor::SparseProjection),
9676            weight(&[1, 2], &[1.0, 1.0]),
9677        );
9678        weights.insert(
9679            layer_id(0, LayerTensor::SparseCompressorGate),
9680            weight(&[2, 2], &[0.3, -0.2, 0.1, 0.4]),
9681        );
9682        weights.insert(
9683            layer_id(0, LayerTensor::SparseCompressorPosition),
9684            weight(&[4, 2], &[0.1, 0.0, -0.1, 0.2, 0.05, -0.05, 0.0, 0.1]),
9685        );
9686        let x: Vec<f32> = (0..tokens * hidden)
9687            .map(|index| ((index % 5) as f32 - 2.0) * 0.3)
9688            .collect();
9689        let q_resid = x.clone();
9690
9691        // top_k 8 / pool 4 = a 2-pool budget, so every causally visible pool selects.
9692        let kpool = KpoolPlan {
9693            pool: 4,
9694            always_select_tail: true,
9695        };
9696        let allowed = kpool_allowed_tokens(
9697            0, 1, 2, 8, &kpool, &weights, &x, &q_resid, tokens, hidden, q_rank,
9698        )
9699        .unwrap();
9700        // Query 7 sees both complete pools; 8 % 4 == 0 leaves no tail.
9701        assert_eq!(allowed[7], (0..8).collect::<Vec<_>>());
9702        // Query 6: pool [4..=7] ends past the query, so only [0..=3] plus tail [4,5,6].
9703        assert_eq!(allowed[6], vec![0, 1, 2, 3, 4, 5, 6]);
9704        // Query 2 precedes any complete pool: tail only.
9705        assert_eq!(allowed[2], vec![0, 1, 2]);
9706
9707        // Without the tail, queries before the first complete pool have no candidates.
9708        let no_tail = KpoolPlan {
9709            pool: 4,
9710            always_select_tail: false,
9711        };
9712        let error = kpool_allowed_tokens(
9713            0, 1, 2, 8, &no_tail, &weights, &x, &q_resid, tokens, hidden, q_rank,
9714        )
9715        .unwrap_err();
9716        assert!(matches!(
9717            error,
9718            ReferenceError::InvalidPlan {
9719                reason: "k-pool selection produced an empty candidate set for a query",
9720                ..
9721            }
9722        ));
9723    }
9724}