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 % (max_ngram - 1) != 0
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 % (max_ngram - 1) != 0
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
5503fn execute_mtp(
5504    plan: &ModelPlan,
5505    weights: &ReferenceWeights,
5506    token_ids: &[u32],
5507    embedding: &[f32],
5508    trunk_hidden: &[f32],
5509    tokens: usize,
5510    hidden: usize,
5511    vocab: usize,
5512    model_output: &[f32],
5513) -> Result<Vec<ReferenceMtpOutput>, ReferenceError> {
5514    if plan.mtp_blocks.is_empty() {
5515        return Ok(Vec::new());
5516    }
5517    let gated = gated_residual_topology(plan)?;
5518    if gated.is_some() && plan.mtp_blocks.len() > 1 {
5519        // The checkpoint has ONE mtp.* namespace (glue + mixer); a second depth would
5520        // alias its tensors.
5521        return Err(ReferenceError::UnsupportedOperation {
5522            layer: None,
5523            operation: "multi-depth gated-residual MTP",
5524        });
5525    }
5526    let mut embedded = vec![0.0; tokens * hidden];
5527    for (position, &token) in token_ids.iter().enumerate() {
5528        let token = token as usize;
5529        embedded[position * hidden..(position + 1) * hidden]
5530            .copy_from_slice(&embedding[token * hidden..(token + 1) * hidden]);
5531    }
5532    let mut source_hidden = trunk_hidden.to_vec();
5533    let mut outputs = Vec::with_capacity(plan.mtp_blocks.len());
5534    for block in &plan.mtp_blocks {
5535        let fused = match block.input.fusion {
5536            memra_gguf::model_plan::MtpFusionPlan::ConcatenateProjection => {
5537                if source_hidden.len() != tokens * hidden {
5538                    return Err(ReferenceError::UnsupportedOperation {
5539                        layer: None,
5540                        operation: "HyperConnections MTP fusion",
5541                    });
5542                }
5543                let embedding_norm = rms_norm(
5544                    &embedded,
5545                    tokens,
5546                    hidden,
5547                    tensor(
5548                        weights,
5549                        &TensorId::Mtp {
5550                            depth: block.depth,
5551                            tensor: MtpTensor::EmbeddingNorm,
5552                        },
5553                        &[hidden],
5554                    )?,
5555                    block.input.embedding_norm.epsilon,
5556                );
5557                let hidden_norm = rms_norm(
5558                    &source_hidden,
5559                    tokens,
5560                    hidden,
5561                    tensor(
5562                        weights,
5563                        &TensorId::Mtp {
5564                            depth: block.depth,
5565                            tensor: MtpTensor::HiddenNorm,
5566                        },
5567                        &[hidden],
5568                    )?,
5569                    block.input.hidden_norm.epsilon,
5570                );
5571                let mut concatenated = vec![0.0; tokens * 2 * hidden];
5572                for token in 0..tokens {
5573                    concatenated[token * 2 * hidden..token * 2 * hidden + hidden]
5574                        .copy_from_slice(&embedding_norm[token * hidden..(token + 1) * hidden]);
5575                    concatenated[token * 2 * hidden + hidden..(token + 1) * 2 * hidden]
5576                        .copy_from_slice(&hidden_norm[token * hidden..(token + 1) * hidden]);
5577                }
5578                linear(
5579                    &concatenated,
5580                    tensor(
5581                        weights,
5582                        &TensorId::Mtp {
5583                            depth: block.depth,
5584                            tensor: MtpTensor::FusionProjection,
5585                        },
5586                        &[hidden, 2 * hidden],
5587                    )?,
5588                    tokens,
5589                    2 * hidden,
5590                    hidden,
5591                )
5592            }
5593            memra_gguf::model_plan::MtpFusionPlan::SeparateProjections => {
5594                // qwen4_exp (SEMANTICS.md §MTP, sglang_qwen4_exp_mtp.py L105-115): the
5595                // draft input is the trunk's WIDE state, normed FLAT over the full wide
5596                // vector (GemmaRMSNorm(hc_count*hidden) — not grouped), viewed per stream
5597                // through fc_hidden, plus fc_embedding(norm(embed)) broadcast over streams.
5598                let Some((streams, _)) = gated else {
5599                    return Err(ReferenceError::InvalidPlan {
5600                        layer: Some(block.layer.index),
5601                        reason: "separate-projection MTP fusion requires a gated-residual trunk",
5602                    });
5603                };
5604                let wide = streams * hidden;
5605                if source_hidden.len() != tokens * wide {
5606                    return Err(ReferenceError::InvalidPlan {
5607                        layer: Some(block.layer.index),
5608                        reason: "separate-projection MTP fusion requires the wide trunk state",
5609                    });
5610                }
5611                let embedding_norm = rms_norm(
5612                    &embedded,
5613                    tokens,
5614                    hidden,
5615                    tensor(
5616                        weights,
5617                        &TensorId::Mtp {
5618                            depth: block.depth,
5619                            tensor: MtpTensor::EmbeddingNorm,
5620                        },
5621                        &[hidden],
5622                    )?,
5623                    block.input.embedding_norm.epsilon,
5624                );
5625                let embedding_projected = linear(
5626                    &embedding_norm,
5627                    tensor(
5628                        weights,
5629                        &TensorId::Mtp {
5630                            depth: block.depth,
5631                            tensor: MtpTensor::EmbeddingProjection,
5632                        },
5633                        &[hidden, hidden],
5634                    )?,
5635                    tokens,
5636                    hidden,
5637                    hidden,
5638                );
5639                let hidden_norm = rms_norm(
5640                    &source_hidden,
5641                    tokens,
5642                    wide,
5643                    tensor(
5644                        weights,
5645                        &TensorId::Mtp {
5646                            depth: block.depth,
5647                            tensor: MtpTensor::HiddenNorm,
5648                        },
5649                        &[wide],
5650                    )?,
5651                    block.input.hidden_norm.epsilon,
5652                );
5653                let hidden_projected = linear(
5654                    &hidden_norm,
5655                    tensor(
5656                        weights,
5657                        &TensorId::Mtp {
5658                            depth: block.depth,
5659                            tensor: MtpTensor::HiddenProjection,
5660                        },
5661                        &[hidden, hidden],
5662                    )?,
5663                    tokens * streams,
5664                    hidden,
5665                    hidden,
5666                );
5667                let mut fused = hidden_projected;
5668                for token in 0..tokens {
5669                    for stream in 0..streams {
5670                        for column in 0..hidden {
5671                            fused[(token * streams + stream) * hidden + column] +=
5672                                embedding_projected[token * hidden + column];
5673                        }
5674                    }
5675                }
5676                fused
5677            }
5678        };
5679        let (hidden_next, state) = execute_layer(
5680            &block.layer,
5681            weights,
5682            &fused,
5683            token_ids,
5684            tokens,
5685            hidden,
5686            vocab,
5687            LayerScope::Mtp { depth: block.depth },
5688        )?;
5689        let norm_id = TensorId::Mtp {
5690            depth: block.depth,
5691            tensor: MtpTensor::OutputNorm,
5692        };
5693        let final_hidden = if let Some((streams, rank)) = gated {
5694            // The draft exits through its OWN hyper_connection_mixer (SEMANTICS.md §MTP);
5695            // there is no MTP final norm and no model OutputNorm to fall back to.
5696            gated_residual_read(
5697                weights,
5698                LayerScope::Mtp { depth: block.depth }.mixer_prefix(),
5699                "",
5700                &hidden_next,
5701                tokens,
5702                streams,
5703                hidden,
5704                rank,
5705                plan.output_norm.epsilon,
5706                false,
5707            )?
5708            .0
5709        } else {
5710            let norm = match weights.get(&norm_id) {
5711                Some(tensor) => tensor_checked(&norm_id, tensor, &[hidden])?,
5712                None => tensor(weights, &TensorId::OutputNorm, &[hidden])?,
5713            };
5714            rms_norm(&hidden_next, tokens, hidden, norm, plan.output_norm.epsilon)
5715        };
5716        let head_id = TensorId::Mtp {
5717            depth: block.depth,
5718            tensor: MtpTensor::OutputProjection,
5719        };
5720        let head = match weights.get(&head_id) {
5721            Some(tensor) => tensor_checked(&head_id, tensor, &[vocab, hidden])?,
5722            None => model_output,
5723        };
5724        let mut logits = linear(&final_hidden, head, tokens, hidden, vocab);
5725        apply_logits_transforms(&mut logits, vocab, &plan.logits);
5726        source_hidden = hidden_next.clone();
5727        outputs.push(ReferenceMtpOutput {
5728            depth: block.depth,
5729            logits,
5730            hidden: hidden_next,
5731            state,
5732        });
5733    }
5734    Ok(outputs)
5735}
5736
5737fn mla_attention(
5738    layer: u32,
5739    plan: &memra_gguf::model_plan::MlaAttentionPlan,
5740    epsilon: f32,
5741    weights: &ReferenceWeights,
5742    x: &[f32],
5743    tokens: usize,
5744    hidden: usize,
5745) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
5746    if let memra_gguf::model_plan::MlaAttentionPlan::CompressedKv { .. } = plan {
5747        return compressed_mla_attention(layer, plan, epsilon, weights, x, tokens, hidden);
5748    }
5749    let memra_gguf::model_plan::MlaAttentionPlan::LatentKv {
5750        query_heads,
5751        q_lora_rank,
5752        kv_lora_rank,
5753        qk_head_dim,
5754        rope_head_dim,
5755        value_head_dim,
5756        rope,
5757        sparse_index,
5758    } = plan.clone()
5759    else {
5760        return Err(ReferenceError::UnsupportedOperation {
5761            layer: Some(layer),
5762            operation: "compressed-KV MLA",
5763        });
5764    };
5765    // Per-token indexers execute only through full-selection equivalence; the k-pool
5766    // indexer (glm5_next) selects for real and is scored after q_resid exists.
5767    let plain_sparse_top_k = match &sparse_index {
5768        memra_gguf::model_plan::SparseIndexPlan::None
5769        | memra_gguf::model_plan::SparseIndexPlan::Own { kpool: Some(_), .. } => None,
5770        memra_gguf::model_plan::SparseIndexPlan::Own {
5771            top_k, kpool: None, ..
5772        }
5773        | memra_gguf::model_plan::SparseIndexPlan::SharedFromPrevious { top_k } => {
5774            Some(*top_k as usize)
5775        }
5776    };
5777    if plain_sparse_top_k.is_some_and(|top_k| tokens > top_k) {
5778        return Err(ReferenceError::UnsupportedOperation {
5779            layer: Some(layer),
5780            operation: "sparse MLA selection beyond full-selection equivalence",
5781        });
5782    }
5783    let heads = query_heads as usize;
5784    let q_rank = q_lora_rank as usize;
5785    let kv_rank = kv_lora_rank as usize;
5786    let qk_dim = qk_head_dim as usize;
5787    let rope_dim = rope_head_dim as usize;
5788    let nope_dim = qk_dim - rope_dim;
5789    let value_dim = value_head_dim as usize;
5790    let latent_dim = kv_rank + rope_dim;
5791
5792    let q_down = linear(
5793        x,
5794        tensor(
5795            weights,
5796            &layer_id(layer, LayerTensor::MlaQueryDown),
5797            &[q_rank, hidden],
5798        )?,
5799        tokens,
5800        hidden,
5801        q_rank,
5802    );
5803    let q_down = rms_norm(
5804        &q_down,
5805        tokens,
5806        q_rank,
5807        tensor(
5808            weights,
5809            &layer_id(layer, LayerTensor::MlaQueryDownNorm),
5810            &[q_rank],
5811        )?,
5812        epsilon,
5813    );
5814    // q_down is q_resid = q_a_layernorm(q_a_proj(x)): it feeds both the MLA query
5815    // up-projection and the k-pool indexer.
5816    let allowed_mask = match &sparse_index {
5817        memra_gguf::model_plan::SparseIndexPlan::Own {
5818            heads: index_heads,
5819            head_dim: index_dim,
5820            top_k,
5821            kpool: Some(kpool),
5822        } => {
5823            let allowed = kpool_allowed_tokens(
5824                layer,
5825                *index_heads as usize,
5826                *index_dim as usize,
5827                *top_k as usize,
5828                kpool,
5829                weights,
5830                x,
5831                &q_down,
5832                tokens,
5833                hidden,
5834                q_rank,
5835            )?;
5836            let mut mask = vec![false; tokens * tokens];
5837            for (token, sources) in allowed.iter().enumerate() {
5838                for &source in sources {
5839                    mask[token * tokens + source] = true;
5840                }
5841            }
5842            Some(mask)
5843        }
5844        _ => None,
5845    };
5846    let query = linear(
5847        &q_down,
5848        tensor(
5849            weights,
5850            &layer_id(layer, LayerTensor::MlaQueryUp),
5851            &[heads * qk_dim, q_rank],
5852        )?,
5853        tokens,
5854        q_rank,
5855        heads * qk_dim,
5856    );
5857    let latent_raw = linear(
5858        x,
5859        tensor(
5860            weights,
5861            &layer_id(layer, LayerTensor::MlaKvDown),
5862            &[latent_dim, hidden],
5863        )?,
5864        tokens,
5865        hidden,
5866        latent_dim,
5867    );
5868    let kv_norm = tensor(
5869        weights,
5870        &layer_id(layer, LayerTensor::MlaKvDownNorm),
5871        &[kv_rank],
5872    )?;
5873    let mut latent = latent_raw;
5874    for token in 0..tokens {
5875        let offset = token * latent_dim;
5876        let normalized = rms_norm(
5877            &latent[offset..offset + kv_rank],
5878            1,
5879            kv_rank,
5880            kv_norm,
5881            epsilon,
5882        );
5883        latent[offset..offset + kv_rank].copy_from_slice(&normalized);
5884    }
5885
5886    let mut query_nope = vec![0.0; tokens * heads * nope_dim];
5887    let mut query_rope = vec![0.0; tokens * heads * rope_dim];
5888    for token in 0..tokens {
5889        for head in 0..heads {
5890            let source = (token * heads + head) * qk_dim;
5891            let nope_target = (token * heads + head) * nope_dim;
5892            let rope_target = (token * heads + head) * rope_dim;
5893            query_nope[nope_target..nope_target + nope_dim]
5894                .copy_from_slice(&query[source..source + nope_dim]);
5895            query_rope[rope_target..rope_target + rope_dim]
5896                .copy_from_slice(&query[source + nope_dim..source + qk_dim]);
5897        }
5898    }
5899    let (rope_factors, rope_mscale) = rope_factor_values(&rope, weights)?;
5900    apply_rope(
5901        &mut query_rope,
5902        tokens,
5903        heads,
5904        rope_dim,
5905        rope.dimensions as usize,
5906        rope.base,
5907        rope_factors.as_deref(),
5908        rope_mscale,
5909    );
5910    let mut key_rope = vec![0.0; tokens * rope_dim];
5911    for token in 0..tokens {
5912        key_rope[token * rope_dim..(token + 1) * rope_dim]
5913            .copy_from_slice(&latent[token * latent_dim + kv_rank..(token + 1) * latent_dim]);
5914    }
5915    apply_rope(
5916        &mut key_rope,
5917        tokens,
5918        1,
5919        rope_dim,
5920        rope.dimensions as usize,
5921        rope.base,
5922        rope_factors.as_deref(),
5923        rope_mscale,
5924    );
5925    for token in 0..tokens {
5926        latent[token * latent_dim + kv_rank..(token + 1) * latent_dim]
5927            .copy_from_slice(&key_rope[token * rope_dim..(token + 1) * rope_dim]);
5928    }
5929
5930    // Contract layout: [head][kv_rank][nope] (see `deterministic_fixture`).
5931    let key_weight = tensor(
5932        weights,
5933        &layer_id(layer, LayerTensor::MlaKeyUp),
5934        &[heads, kv_rank, nope_dim],
5935    )?;
5936    let value_weight = tensor(
5937        weights,
5938        &layer_id(layer, LayerTensor::MlaValueUp),
5939        &[heads, value_dim, kv_rank],
5940    )?;
5941    let mut key_nope = vec![0.0; tokens * heads * nope_dim];
5942    let mut value = vec![0.0; tokens * heads * value_dim];
5943    for token in 0..tokens {
5944        let latent_row = &latent[token * latent_dim..token * latent_dim + kv_rank];
5945        for head in 0..heads {
5946            for out in 0..nope_dim {
5947                for rank in 0..kv_rank {
5948                    key_nope[(token * heads + head) * nope_dim + out] +=
5949                        latent_row[rank] * key_weight[(head * kv_rank + rank) * nope_dim + out];
5950                }
5951            }
5952            for out in 0..value_dim {
5953                for rank in 0..kv_rank {
5954                    value[(token * heads + head) * value_dim + out] +=
5955                        latent_row[rank] * value_weight[(head * value_dim + out) * kv_rank + rank];
5956                }
5957            }
5958        }
5959    }
5960    let mut attended = vec![0.0; tokens * heads * value_dim];
5961    let scale = 1.0 / (qk_dim as f32).sqrt();
5962    for token in 0..tokens {
5963        for head in 0..heads {
5964            let mut scores = Vec::with_capacity(token + 1);
5965            for source in 0..=token {
5966                // The indexer's allowed set masks keys exactly like the eager
5967                // additive -inf mask built from topk_indices.
5968                if allowed_mask
5969                    .as_ref()
5970                    .is_some_and(|mask| !mask[token * tokens + source])
5971                {
5972                    scores.push(f32::NEG_INFINITY);
5973                    continue;
5974                }
5975                let mut score = 0.0;
5976                for dim in 0..nope_dim {
5977                    score += query_nope[(token * heads + head) * nope_dim + dim]
5978                        * key_nope[(source * heads + head) * nope_dim + dim];
5979                }
5980                for dim in 0..rope_dim {
5981                    score += query_rope[(token * heads + head) * rope_dim + dim]
5982                        * key_rope[source * rope_dim + dim];
5983                }
5984                scores.push(score * scale);
5985            }
5986            softmax_in_place(&mut scores);
5987            for (source, probability) in scores.into_iter().enumerate() {
5988                for dim in 0..value_dim {
5989                    attended[(token * heads + head) * value_dim + dim] +=
5990                        probability * value[(source * heads + head) * value_dim + dim];
5991                }
5992            }
5993        }
5994    }
5995    let output = linear(
5996        &attended,
5997        tensor(
5998            weights,
5999            &layer_id(layer, LayerTensor::MlaOutput),
6000            &[hidden, heads * value_dim],
6001        )?,
6002        tokens,
6003        heads * value_dim,
6004        hidden,
6005    );
6006    Ok((
6007        output,
6008        ReferenceLayerState::LatentKv {
6009            rows: latent,
6010            tokens,
6011            width: latent_dim,
6012        },
6013    ))
6014}
6015
6016/// K-pool compressed indexer selection (Glm5NextTextIndexer.forward), single-sequence
6017/// causal case: every token is a valid key, so pooling starts at index 0 and only
6018/// causality masks candidates. Returns the allowed source-token set per query,
6019/// sorted ascending.
6020///
6021/// PUBLIC because it is the CUDA k-pool indexer's oracle: `memra-engine`'s
6022/// `tests/glm5_kpool_indexer_gpu.rs` compares the device's selected index sets against this
6023/// function on identical inputs. Its scope line above is part of the contract — a padded or
6024/// batched caller is outside it.
6025#[allow(clippy::too_many_arguments)]
6026pub fn kpool_allowed_tokens(
6027    layer: u32,
6028    index_heads: usize,
6029    index_dim: usize,
6030    top_k: usize,
6031    kpool: &memra_gguf::model_plan::KpoolPlan,
6032    weights: &ReferenceWeights,
6033    x: &[f32],
6034    q_resid: &[f32],
6035    tokens: usize,
6036    hidden: usize,
6037    q_rank: usize,
6038) -> Result<Vec<Vec<usize>>, ReferenceError> {
6039    let pool = kpool.pool as usize;
6040    if index_heads == 0 || index_dim == 0 || pool == 0 {
6041        return Err(ReferenceError::InvalidPlan {
6042            layer: Some(layer),
6043            reason: "k-pool sparse index requires positive heads, head_dim, and pool",
6044        });
6045    }
6046    let q = linear(
6047        q_resid,
6048        tensor(
6049            weights,
6050            &layer_id(layer, LayerTensor::SparseQuery),
6051            &[index_heads * index_dim, q_rank],
6052        )?,
6053        tokens,
6054        q_rank,
6055        index_heads * index_dim,
6056    );
6057    let key = layer_norm(
6058        &linear(
6059            x,
6060            tensor(
6061                weights,
6062                &layer_id(layer, LayerTensor::SparseKey),
6063                &[index_dim, hidden],
6064            )?,
6065            tokens,
6066            hidden,
6067            index_dim,
6068        ),
6069        tokens,
6070        index_dim,
6071        tensor(
6072            weights,
6073            &layer_id(layer, LayerTensor::SparseKeyNorm),
6074            &[index_dim],
6075        )?,
6076        tensor(
6077            weights,
6078            &layer_id(layer, LayerTensor::SparseKeyNormBias),
6079            &[index_dim],
6080        )?,
6081    );
6082    let gate_scores = linear(
6083        x,
6084        tensor(
6085            weights,
6086            &layer_id(layer, LayerTensor::SparseCompressorGate),
6087            &[index_dim, hidden],
6088        )?,
6089        tokens,
6090        hidden,
6091        index_dim,
6092    );
6093    let ape = tensor(
6094        weights,
6095        &layer_id(layer, LayerTensor::SparseCompressorPosition),
6096        &[pool, index_dim],
6097    )?;
6098    // Only COMPLETE pools are candidates; the incomplete tail never scores. Each
6099    // channel takes its own softmax over the pool members (gate score + APE).
6100    let pools = tokens / pool;
6101    let mut pool_keys = vec![0.0f32; pools * index_dim];
6102    for pool_index in 0..pools {
6103        for channel in 0..index_dim {
6104            let mut logits = Vec::with_capacity(pool);
6105            for slot in 0..pool {
6106                logits.push(
6107                    gate_scores[(pool_index * pool + slot) * index_dim + channel]
6108                        + ape[slot * index_dim + channel],
6109                );
6110            }
6111            softmax_in_place(&mut logits);
6112            let mut pooled = 0.0;
6113            for slot in 0..pool {
6114                pooled += logits[slot] * key[(pool_index * pool + slot) * index_dim + channel];
6115            }
6116            pool_keys[pool_index * index_dim + channel] = pooled;
6117        }
6118    }
6119    let mut head_weights = linear(
6120        x,
6121        tensor(
6122            weights,
6123            &layer_id(layer, LayerTensor::SparseProjection),
6124            &[index_heads, hidden],
6125        )?,
6126        tokens,
6127        hidden,
6128        index_heads,
6129    );
6130    let head_scale = (index_heads as f32).powf(-0.5);
6131    for value in &mut head_weights {
6132        *value *= head_scale;
6133    }
6134    // Same scale convention as the per-token DSA indexer: relu(q . k * hd^-0.5).
6135    let softmax_scale = (index_dim as f32).powf(-0.5);
6136    let select_k = (top_k / pool).min(pools);
6137    let mut allowed = Vec::with_capacity(tokens);
6138    for token in 0..tokens {
6139        // A pool is selectable only when its final token index is <= the query.
6140        let visible_pools = ((token + 1) / pool).min(pools);
6141        let mut scored: Vec<(usize, f32)> = (0..visible_pools)
6142            .map(|pool_index| {
6143                let mut score = 0.0f32;
6144                for head in 0..index_heads {
6145                    let mut dot = 0.0f32;
6146                    for dim in 0..index_dim {
6147                        dot += q[(token * index_heads + head) * index_dim + dim]
6148                            * pool_keys[pool_index * index_dim + dim];
6149                    }
6150                    score +=
6151                        (dot * softmax_scale).max(0.0) * head_weights[token * index_heads + head];
6152                }
6153                (pool_index, score)
6154            })
6155            .collect();
6156        scored.sort_by(|left, right| {
6157            right
6158                .1
6159                .partial_cmp(&left.1)
6160                .unwrap_or(std::cmp::Ordering::Equal)
6161                .then(left.0.cmp(&right.0))
6162        });
6163        let mut selected: Vec<usize> = Vec::new();
6164        for &(pool_index, _) in scored.iter().take(select_k) {
6165            selected.extend(pool_index * pool..(pool_index + 1) * pool);
6166        }
6167        if kpool.always_select_tail {
6168            // The current incomplete tail: the visible tokens past the last complete
6169            // visible pool (at most pool - 1 of them), always <= the query index.
6170            let visible = token + 1;
6171            let tail = visible % pool;
6172            selected.extend(visible - tail..visible);
6173        }
6174        if selected.is_empty() {
6175            // always_select_tail=false leaves early queries (before the first
6176            // complete pool) with no candidates; the reference would emit NaN rows.
6177            return Err(ReferenceError::InvalidPlan {
6178                layer: Some(layer),
6179                reason: "k-pool selection produced an empty candidate set for a query",
6180            });
6181        }
6182        selected.sort_unstable();
6183        allowed.push(selected);
6184    }
6185    Ok(allowed)
6186}
6187
6188#[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
6189fn compressed_mla_attention(
6190    layer: u32,
6191    plan: &memra_gguf::model_plan::MlaAttentionPlan,
6192    epsilon: f32,
6193    weights: &ReferenceWeights,
6194    x: &[f32],
6195    tokens: usize,
6196    hidden: usize,
6197) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
6198    use memra_gguf::dsv4_forward::{
6199        ActQuantVariant, IndexerW, apply_rope as apply_dsv4_rope, matmul, precompute_freqs_cis,
6200        rmsnorm,
6201    };
6202    use memra_gguf::model_plan::{MlaAttentionPlan, RopeFactors, SparseIndexPlan};
6203
6204    let MlaAttentionPlan::CompressedKv {
6205        query_heads,
6206        q_lora_rank,
6207        latent_head_dim,
6208        rope_head_dim,
6209        output_lora_rank,
6210        output_groups,
6211        window,
6212        rope,
6213        compressor,
6214        sparse_index,
6215    } = plan
6216    else {
6217        unreachable!()
6218    };
6219    let heads = *query_heads as usize;
6220    let q_rank = *q_lora_rank as usize;
6221    let head_dim = *latent_head_dim as usize;
6222    let rope_dim = *rope_head_dim as usize;
6223    let output_rank = *output_lora_rank as usize;
6224    let groups = *output_groups as usize;
6225    let window = *window as usize;
6226    if heads == 0
6227        || q_rank == 0
6228        || head_dim == 0
6229        || rope_dim == 0
6230        || rope_dim > head_dim
6231        || !(head_dim - rope_dim).is_multiple_of(64)
6232        || groups == 0
6233        || heads % groups != 0
6234        || window == 0
6235    {
6236        return Err(ReferenceError::InvalidPlan {
6237            layer: Some(layer),
6238            reason: "compressed attention has invalid reference geometry",
6239        });
6240    }
6241    let (original_context, factor, beta_fast, beta_slow) = match rope.factors {
6242        RopeFactors::None => (0, 1.0, 32.0, 1.0),
6243        RopeFactors::Yarn {
6244            factor,
6245            original_context,
6246            beta_fast,
6247            beta_slow,
6248        } => (original_context, factor, beta_fast, beta_slow),
6249        _ => {
6250            return Err(ReferenceError::InvalidPlan {
6251                layer: Some(layer),
6252                reason: "compressed attention requires plain or YaRN RoPE",
6253            });
6254        }
6255    };
6256    let frequencies = precompute_freqs_cis(
6257        rope_dim,
6258        tokens.max(1),
6259        original_context,
6260        rope.base,
6261        factor,
6262        beta_fast,
6263        beta_slow,
6264    );
6265    let positions: Vec<usize> = (0..tokens).collect();
6266
6267    let query_low_rank = rmsnorm(
6268        &matmul(
6269            x,
6270            tokens,
6271            hidden,
6272            tensor(
6273                weights,
6274                &layer_id(layer, LayerTensor::MlaQueryDown),
6275                &[q_rank, hidden],
6276            )?,
6277            q_rank,
6278        ),
6279        tensor(
6280            weights,
6281            &layer_id(layer, LayerTensor::MlaQueryDownNorm),
6282            &[q_rank],
6283        )?,
6284        epsilon,
6285    );
6286    let mut query = matmul(
6287        &query_low_rank,
6288        tokens,
6289        q_rank,
6290        tensor(
6291            weights,
6292            &layer_id(layer, LayerTensor::MlaQueryUp),
6293            &[heads * head_dim, q_rank],
6294        )?,
6295        heads * head_dim,
6296    );
6297    for head in query.chunks_exact_mut(head_dim) {
6298        let mean_square = head
6299            .iter()
6300            .map(|value| (*value as f64) * (*value as f64))
6301            .sum::<f64>()
6302            / head_dim as f64;
6303        let scale = 1.0 / (mean_square as f32 + epsilon).sqrt();
6304        for value in head {
6305            *value *= scale;
6306        }
6307    }
6308    apply_dsv4_rope(
6309        &mut query,
6310        tokens,
6311        heads,
6312        head_dim,
6313        rope_dim,
6314        &frequencies,
6315        &positions,
6316        false,
6317    );
6318
6319    let mut key_value = rmsnorm(
6320        &matmul(
6321            x,
6322            tokens,
6323            hidden,
6324            tensor(
6325                weights,
6326                &layer_id(layer, LayerTensor::MlaKvDown),
6327                &[head_dim, hidden],
6328            )?,
6329            head_dim,
6330        ),
6331        tensor(
6332            weights,
6333            &layer_id(layer, LayerTensor::MlaKvDownNorm),
6334            &[head_dim],
6335        )?,
6336        epsilon,
6337    );
6338    apply_dsv4_rope(
6339        &mut key_value,
6340        tokens,
6341        1,
6342        head_dim,
6343        rope_dim,
6344        &frequencies,
6345        &positions,
6346        false,
6347    );
6348    for row in key_value.chunks_exact_mut(head_dim) {
6349        memra_gguf::dsv4_forward::act_quant(
6350            &mut row[..head_dim - rope_dim],
6351            64,
6352            ActQuantVariant::RefFp8Round,
6353        );
6354    }
6355
6356    let (mut indices, mut slots) = memra_gguf::dsv4_forward::window_topk_idxs(window, tokens);
6357    let mut key_value_rows = tokens;
6358    let mut compressed_tokens = 0;
6359    if let Some(compressor_plan) = compressor {
6360        let ratio = compressor_plan.ratio as usize;
6361        let compressor = reference_compressor(
6362            weights,
6363            layer,
6364            hidden,
6365            head_dim,
6366            ratio,
6367            compressor_plan.latent_dim as usize,
6368            false,
6369        )?;
6370        let (compressed_indices, compressed_slots) = match sparse_index {
6371            SparseIndexPlan::None => {
6372                memra_gguf::dsv4_forward::compress_topk_idxs(ratio, tokens, tokens)
6373            }
6374            SparseIndexPlan::Own {
6375                heads: index_heads,
6376                head_dim: index_dim,
6377                top_k,
6378                kpool,
6379            } => {
6380                // K-pool scoring is a LatentKv (glm5_next) program; dsv4 compiles None.
6381                if kpool.is_some() {
6382                    return Err(ReferenceError::UnsupportedOperation {
6383                        layer: Some(layer),
6384                        operation: "k-pool sparse index on compressed attention",
6385                    });
6386                }
6387                let index_heads = *index_heads as usize;
6388                let index_dim = *index_dim as usize;
6389                if index_dim < rope_dim
6390                    || !index_dim.is_multiple_of(32)
6391                    || !index_dim.is_power_of_two()
6392                {
6393                    return Err(ReferenceError::InvalidPlan {
6394                        layer: Some(layer),
6395                        reason: "compressed sparse index has invalid head geometry",
6396                    });
6397                }
6398                let indexer = IndexerW {
6399                    wq_b: tensor(
6400                        weights,
6401                        &layer_id(layer, LayerTensor::SparseQuery),
6402                        &[index_heads * index_dim, q_rank],
6403                    )?
6404                    .to_vec(),
6405                    weights_proj: tensor(
6406                        weights,
6407                        &layer_id(layer, LayerTensor::SparseProjection),
6408                        &[index_heads, hidden],
6409                    )?
6410                    .to_vec(),
6411                    compressor: reference_compressor(
6412                        weights,
6413                        layer,
6414                        hidden,
6415                        index_dim,
6416                        ratio,
6417                        2 * index_dim,
6418                        true,
6419                    )?,
6420                    heads: index_heads,
6421                    hd: index_dim,
6422                    topk: *top_k as usize,
6423                };
6424                let output = indexer.forward(
6425                    x,
6426                    &query_low_rank,
6427                    tokens,
6428                    hidden,
6429                    q_rank,
6430                    tokens,
6431                    &frequencies,
6432                    rope_dim,
6433                    epsilon,
6434                    ActQuantVariant::RefFp8Round,
6435                    false,
6436                );
6437                (output.idxs, output.slots)
6438            }
6439            SparseIndexPlan::SharedFromPrevious { .. } => {
6440                return Err(ReferenceError::UnsupportedOperation {
6441                    layer: Some(layer),
6442                    operation: "shared compressed sparse-index execution",
6443                });
6444            }
6445        };
6446        if compressed_slots > 0 {
6447            let mut merged = vec![-1; tokens * (slots + compressed_slots)];
6448            for token in 0..tokens {
6449                merged[token * (slots + compressed_slots)
6450                    ..token * (slots + compressed_slots) + slots]
6451                    .copy_from_slice(&indices[token * slots..(token + 1) * slots]);
6452                merged[token * (slots + compressed_slots) + slots
6453                    ..(token + 1) * (slots + compressed_slots)]
6454                    .copy_from_slice(
6455                        &compressed_indices
6456                            [token * compressed_slots..(token + 1) * compressed_slots],
6457                    );
6458            }
6459            indices = merged;
6460            slots += compressed_slots;
6461        }
6462        if let Some((compressed, count)) = compressor.forward(
6463            x,
6464            tokens,
6465            hidden,
6466            &frequencies,
6467            rope_dim,
6468            epsilon,
6469            ActQuantVariant::RefFp8Round,
6470        ) {
6471            key_value.extend_from_slice(&compressed);
6472            key_value_rows += count;
6473            compressed_tokens = count;
6474        }
6475    }
6476
6477    let sink = tensor(
6478        weights,
6479        &layer_id(layer, LayerTensor::AttentionSink),
6480        &[heads],
6481    )?;
6482    let attention_scale = (head_dim as f64).powf(-0.5) as f32;
6483    let mut attended = vec![0.0; tokens * heads * head_dim];
6484    for token in 0..tokens {
6485        let selected = &indices[token * slots..(token + 1) * slots];
6486        memra_gguf::dsv4_decode::sparse_attn_query(
6487            &query[token * heads * head_dim..(token + 1) * heads * head_dim],
6488            heads,
6489            head_dim,
6490            selected,
6491            |index| &key_value[index * head_dim..(index + 1) * head_dim],
6492            sink,
6493            attention_scale,
6494            &mut attended[token * heads * head_dim..(token + 1) * heads * head_dim],
6495        );
6496    }
6497    apply_dsv4_rope(
6498        &mut attended,
6499        tokens,
6500        heads,
6501        head_dim,
6502        rope_dim,
6503        &frequencies,
6504        &positions,
6505        true,
6506    );
6507
6508    let group_width = heads / groups * head_dim;
6509    let output_down = tensor(
6510        weights,
6511        &layer_id(layer, LayerTensor::MlaOutputDown),
6512        &[groups * output_rank, group_width],
6513    )?;
6514    let mut grouped = vec![0.0; tokens * groups * output_rank];
6515    for token in 0..tokens {
6516        for group in 0..groups {
6517            let source = &attended[token * heads * head_dim + group * group_width
6518                ..token * heads * head_dim + (group + 1) * group_width];
6519            let group_weight = &output_down
6520                [group * output_rank * group_width..(group + 1) * output_rank * group_width];
6521            for rank in 0..output_rank {
6522                grouped[(token * groups + group) * output_rank + rank] =
6523                    memra_gguf::dsv4_forward::dot(
6524                        source,
6525                        &group_weight[rank * group_width..(rank + 1) * group_width],
6526                    );
6527            }
6528        }
6529    }
6530    let output = matmul(
6531        &grouped,
6532        tokens,
6533        groups * output_rank,
6534        tensor(
6535            weights,
6536            &layer_id(layer, LayerTensor::MlaOutput),
6537            &[hidden, groups * output_rank],
6538        )?,
6539        hidden,
6540    );
6541    Ok((
6542        output,
6543        ReferenceLayerState::CompressedAttention {
6544            rows: key_value,
6545            tokens: key_value_rows,
6546            width: head_dim,
6547            window,
6548            compressed_tokens,
6549        },
6550    ))
6551}
6552
6553#[allow(clippy::too_many_arguments)]
6554fn reference_compressor(
6555    weights: &ReferenceWeights,
6556    layer: u32,
6557    hidden: usize,
6558    output_dim: usize,
6559    ratio: usize,
6560    latent: usize,
6561    sparse: bool,
6562) -> Result<memra_gguf::dsv4_forward::CompressorW, ReferenceError> {
6563    let (key_value, gate, norm, position) = if sparse {
6564        (
6565            LayerTensor::SparseCompressorKeyValue,
6566            LayerTensor::SparseCompressorGate,
6567            LayerTensor::SparseCompressorNorm,
6568            LayerTensor::SparseCompressorPosition,
6569        )
6570    } else {
6571        (
6572            LayerTensor::KvCompressorKeyValue,
6573            LayerTensor::KvCompressorGate,
6574            LayerTensor::KvCompressorNorm,
6575            LayerTensor::KvCompressorPosition,
6576        )
6577    };
6578    Ok(memra_gguf::dsv4_forward::CompressorW {
6579        ratio,
6580        d: output_dim,
6581        latent,
6582        overlap: ratio == 4,
6583        rotate: sparse,
6584        wkv: tensor(weights, &layer_id(layer, key_value), &[latent, hidden])?.to_vec(),
6585        wgate: tensor(weights, &layer_id(layer, gate), &[latent, hidden])?.to_vec(),
6586        norm_w: tensor(weights, &layer_id(layer, norm), &[output_dim])?.to_vec(),
6587        ape: tensor(weights, &layer_id(layer, position), &[ratio, latent])?.to_vec(),
6588    })
6589}
6590
6591fn gated_delta_net(
6592    layer: u32,
6593    plan: &memra_gguf::model_plan::GatedDeltaNetPlan,
6594    epsilon: f32,
6595    weights: &ReferenceWeights,
6596    x: &[f32],
6597    tokens: usize,
6598    hidden: usize,
6599) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
6600    let key_heads = plan.key_heads as usize;
6601    let value_heads = plan.value_heads as usize;
6602    let key_dim = plan.key_head_dim as usize;
6603    let value_dim = plan.value_head_dim as usize;
6604    let kernel = plan.conv_kernel as usize;
6605    if key_heads == 0 || value_heads == 0 || key_dim == 0 || value_dim == 0 || kernel == 0 {
6606        return Err(ReferenceError::InvalidPlan {
6607            layer: Some(layer),
6608            reason: "GDN dimensions must be positive",
6609        });
6610    }
6611    let key_width = key_heads * key_dim;
6612    let value_width = value_heads * value_dim;
6613    let conv_width = 2 * key_width + value_width;
6614    let qkv = linear(
6615        x,
6616        tensor(
6617            weights,
6618            &layer_id(layer, LayerTensor::GdnQkv),
6619            &[conv_width, hidden],
6620        )?,
6621        tokens,
6622        hidden,
6623        conv_width,
6624    );
6625    let gate = linear(
6626        x,
6627        tensor(
6628            weights,
6629            &layer_id(layer, LayerTensor::GdnGate),
6630            &[value_width, hidden],
6631        )?,
6632        tokens,
6633        hidden,
6634        value_width,
6635    );
6636    let beta_raw = linear(
6637        x,
6638        tensor(
6639            weights,
6640            &layer_id(layer, LayerTensor::GdnBeta),
6641            &[value_heads, hidden],
6642        )?,
6643        tokens,
6644        hidden,
6645        value_heads,
6646    );
6647    let alpha = linear(
6648        x,
6649        tensor(
6650            weights,
6651            &layer_id(layer, LayerTensor::GdnAlpha),
6652            &[value_heads, hidden],
6653        )?,
6654        tokens,
6655        hidden,
6656        value_heads,
6657    );
6658    let conv_weight = tensor(
6659        weights,
6660        &layer_id(layer, LayerTensor::GdnConv1d),
6661        &[conv_width, kernel],
6662    )?;
6663    let mut conv = vec![0.0; tokens * conv_width];
6664    let pad = kernel - 1;
6665    for token in 0..tokens {
6666        for channel in 0..conv_width {
6667            let mut sum = 0.0;
6668            for tap in 0..kernel {
6669                let source = token as isize - pad as isize + tap as isize;
6670                if source >= 0 {
6671                    sum += qkv[source as usize * conv_width + channel]
6672                        * conv_weight[channel * kernel + tap];
6673                }
6674            }
6675            conv[token * conv_width + channel] = silu(sum);
6676        }
6677    }
6678
6679    let mut query = vec![0.0; tokens * value_heads * key_dim];
6680    let mut key = vec![0.0; tokens * value_heads * key_dim];
6681    let mut value = vec![0.0; tokens * value_width];
6682    for token in 0..tokens {
6683        for value_head in 0..value_heads {
6684            let key_head = value_head % key_heads;
6685            let q_source = token * conv_width + key_head * key_dim;
6686            let k_source = token * conv_width + key_width + key_head * key_dim;
6687            let v_source = token * conv_width + 2 * key_width + value_head * value_dim;
6688            let q_target = (token * value_heads + value_head) * key_dim;
6689            let v_target = (token * value_heads + value_head) * value_dim;
6690            query[q_target..q_target + key_dim]
6691                .copy_from_slice(&conv[q_source..q_source + key_dim]);
6692            key[q_target..q_target + key_dim].copy_from_slice(&conv[k_source..k_source + key_dim]);
6693            value[v_target..v_target + value_dim]
6694                .copy_from_slice(&conv[v_source..v_source + value_dim]);
6695        }
6696    }
6697    l2_normalize_rows(&mut query, tokens * value_heads, key_dim, epsilon);
6698    l2_normalize_rows(&mut key, tokens * value_heads, key_dim, epsilon);
6699
6700    let a = tensor(weights, &layer_id(layer, LayerTensor::GdnA), &[value_heads])?;
6701    let dt = tensor(
6702        weights,
6703        &layer_id(layer, LayerTensor::GdnDtBias),
6704        &[value_heads],
6705    )?;
6706    let mut matrix = vec![0.0; value_heads * value_dim * key_dim];
6707    let mut mixed = vec![0.0; tokens * value_width];
6708    let scale = 1.0 / (key_dim as f32).sqrt();
6709    for token in 0..tokens {
6710        for head in 0..value_heads {
6711            let beta = sigmoid(beta_raw[token * value_heads + head]);
6712            let decay = (a[head] * softplus(alpha[token * value_heads + head] + dt[head])).exp();
6713            let q_offset = (token * value_heads + head) * key_dim;
6714            let v_offset = (token * value_heads + head) * value_dim;
6715            let state_offset = head * value_dim * key_dim;
6716            let mut next = matrix[state_offset..state_offset + value_dim * key_dim].to_vec();
6717            for value_index in 0..value_dim {
6718                let row = state_offset + value_index * key_dim;
6719                let mut state_key = 0.0;
6720                for key_index in 0..key_dim {
6721                    state_key += matrix[row + key_index] * key[q_offset + key_index];
6722                }
6723                let delta = (value[v_offset + value_index] - decay * state_key) * beta;
6724                let mut attended = 0.0;
6725                for key_index in 0..key_dim {
6726                    let updated =
6727                        decay * matrix[row + key_index] + key[q_offset + key_index] * delta;
6728                    next[value_index * key_dim + key_index] = updated;
6729                    attended += updated * query[q_offset + key_index];
6730                }
6731                mixed[v_offset + value_index] = attended * scale;
6732            }
6733            matrix[state_offset..state_offset + value_dim * key_dim].copy_from_slice(&next);
6734        }
6735    }
6736
6737    let norm = tensor(
6738        weights,
6739        &layer_id(layer, LayerTensor::GdnNorm),
6740        &[value_dim],
6741    )?;
6742    let normalized = rms_norm(&mixed, tokens * value_heads, value_dim, norm, epsilon);
6743    let mut gated = normalized;
6744    for index in 0..gated.len() {
6745        // qwen4_exp declares sigmoid here (config output_gate_type) — the ONE numeric
6746        // divergence from the qwen3_5 GDN program (SEMANTICS.md §GDN); every other
6747        // family is the silu arm.
6748        gated[index] *= match plan.gate_activation {
6749            GdnGateActivation::Silu => silu(gate[index]),
6750            GdnGateActivation::Sigmoid => sigmoid(gate[index]),
6751        };
6752    }
6753    let output = linear(
6754        &gated,
6755        tensor(
6756            weights,
6757            &layer_id(layer, LayerTensor::GdnOutput),
6758            &[hidden, value_width],
6759        )?,
6760        tokens,
6761        value_width,
6762        hidden,
6763    );
6764    let mut conv_state = vec![0.0; conv_width * pad];
6765    for channel in 0..conv_width {
6766        for index in 0..pad {
6767            let source = tokens as isize - pad as isize + index as isize;
6768            if source >= 0 {
6769                conv_state[channel * pad + index] = qkv[source as usize * conv_width + channel];
6770            }
6771        }
6772    }
6773    Ok((
6774        output,
6775        ReferenceLayerState::Recurrent {
6776            conv: conv_state,
6777            matrix,
6778            value_heads,
6779            key_head_dim: key_dim,
6780            value_head_dim: value_dim,
6781            conv_width,
6782        },
6783    ))
6784}
6785
6786#[allow(clippy::too_many_arguments)]
6787/// Kimi Delta Attention (recurrent_kimi_delta_attention + Glm5NextTextLinearAttention),
6788/// all f32, sequential over tokens. Only the lower-bound forget-gate branch exists:
6789/// GLM-5.3-Flash always configures `gate_lower_bound`, so the softplus branch of
6790/// Glm5NextTextForgetGate is dead for this model and deliberately not implemented.
6791/// GPU-parity seam: run ONE KDA layer's mixer over `x` (`[tokens, hidden]`, already
6792/// pre-attention-normed) and return its output plus the recurrent state it leaves behind.
6793///
6794/// This is the very `kimi_delta_net` the trunk executor dispatches — exposed so
6795/// `crates/memra-engine/tests/kda_fixture_gpu.rs` can gate the CUDA mixer against the pinned
6796/// reference without standing up a whole model (glm5_next's residual topology and MLA layers
6797/// are a different surface, and a mixer gate must not depend on them).
6798pub fn kimi_delta_net_layer(
6799    layer: u32,
6800    plan: &memra_gguf::model_plan::KimiDeltaNetPlan,
6801    epsilon: f32,
6802    weights: &ReferenceWeights,
6803    x: &[f32],
6804    tokens: usize,
6805    hidden: usize,
6806) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
6807    kimi_delta_net(layer, plan, epsilon, weights, x, tokens, hidden)
6808}
6809
6810fn kimi_delta_net(
6811    layer: u32,
6812    plan: &memra_gguf::model_plan::KimiDeltaNetPlan,
6813    epsilon: f32,
6814    weights: &ReferenceWeights,
6815    x: &[f32],
6816    tokens: usize,
6817    hidden: usize,
6818) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
6819    let heads = plan.num_heads as usize;
6820    let head_dim = plan.head_dim as usize;
6821    let kernel = plan.conv_kernel as usize;
6822    if heads == 0 || head_dim == 0 || kernel == 0 {
6823        return Err(ReferenceError::InvalidPlan {
6824            layer: Some(layer),
6825            reason: "KDA dimensions must be positive",
6826        });
6827    }
6828    let qkv = heads * head_dim;
6829    let conv_width = 3 * qkv;
6830    let project_and_convolve = |projection: LayerTensor,
6831                                conv: LayerTensor|
6832     -> Result<(Vec<f32>, Vec<f32>), ReferenceError> {
6833        let projected = linear(
6834            x,
6835            tensor(weights, &layer_id(layer, projection), &[qkv, hidden])?,
6836            tokens,
6837            hidden,
6838            qkv,
6839        );
6840        // The checkpoint splits the grouped causal conv into three per-plane
6841        // convs; applying each to its own plane is the fused conv exactly.
6842        let conv_weight = tensor(weights, &layer_id(layer, conv), &[qkv, kernel])?;
6843        let mut convolved = vec![0.0; tokens * qkv];
6844        for token in 0..tokens {
6845            for channel in 0..qkv {
6846                let mut sum = 0.0;
6847                for tap in 0..kernel {
6848                    let source = token as isize - (kernel - 1) as isize + tap as isize;
6849                    if source >= 0 {
6850                        sum += projected[source as usize * qkv + channel]
6851                            * conv_weight[channel * kernel + tap];
6852                    }
6853                }
6854                convolved[token * qkv + channel] = silu(sum);
6855            }
6856        }
6857        Ok((projected, convolved))
6858    };
6859    let (q_raw, mut query) =
6860        project_and_convolve(LayerTensor::KdaQuery, LayerTensor::KdaQueryConv)?;
6861    let (k_raw, mut key) = project_and_convolve(LayerTensor::KdaKey, LayerTensor::KdaKeyConv)?;
6862    let (v_raw, value) = project_and_convolve(LayerTensor::KdaValue, LayerTensor::KdaValueConv)?;
6863    // FLA l2norm: x / sqrt(sum(x^2) + 1e-6) — the epsilon sits INSIDE the sqrt and is
6864    // fixed at 1e-6, independent of the layer epsilon.
6865    l2_normalize_rows(&mut query, tokens * heads, head_dim, 1e-6);
6866    l2_normalize_rows(&mut key, tokens * heads, head_dim, 1e-6);
6867    // Query scale head_dim^-0.5 applies AFTER the l2norm.
6868    let query_scale = 1.0 / (head_dim as f32).sqrt();
6869    for entry in &mut query {
6870        *entry *= query_scale;
6871    }
6872
6873    // Forget gate: g = gate_lower_bound * sigmoid(exp(A_log[head]) * (f_b(f_a(x)) + dt_bias)),
6874    // per channel (dt_bias has width qkv).
6875    let forget_down = linear(
6876        x,
6877        tensor(
6878            weights,
6879            &layer_id(layer, LayerTensor::KdaForgetDown),
6880            &[head_dim, hidden],
6881        )?,
6882        tokens,
6883        hidden,
6884        head_dim,
6885    );
6886    let mut forget = linear(
6887        &forget_down,
6888        tensor(
6889            weights,
6890            &layer_id(layer, LayerTensor::KdaForgetUp),
6891            &[qkv, head_dim],
6892        )?,
6893        tokens,
6894        head_dim,
6895        qkv,
6896    );
6897    let dt_bias = tensor(weights, &layer_id(layer, LayerTensor::KdaDtBias), &[qkv])?;
6898    let a_log = tensor(weights, &layer_id(layer, LayerTensor::KdaALog), &[heads])?;
6899    for token in 0..tokens {
6900        #[allow(clippy::needless_range_loop)]
6901        // allow: the explicit index loop keeps the offset arithmetic visible and aligned with the device-side indexing
6902        for head in 0..heads {
6903            let decay_rate = a_log[head].exp();
6904            for dim in 0..head_dim {
6905                let channel = head * head_dim + dim;
6906                let raw = forget[token * qkv + channel] + dt_bias[channel];
6907                forget[token * qkv + channel] = plan.gate_lower_bound * sigmoid(decay_rate * raw);
6908            }
6909        }
6910    }
6911    let beta_raw = linear(
6912        x,
6913        tensor(
6914            weights,
6915            &layer_id(layer, LayerTensor::KdaBeta),
6916            &[heads, hidden],
6917        )?,
6918        tokens,
6919        hidden,
6920        heads,
6921    );
6922
6923    // Recurrence (recurrent_kimi_delta_attention:477-489): state [heads, k_dim, v_dim];
6924    // exp(g) decays along the K dimension.
6925    let mut matrix = vec![0.0; heads * head_dim * head_dim];
6926    let mut core = vec![0.0; tokens * qkv];
6927    for token in 0..tokens {
6928        for head in 0..heads {
6929            let beta = sigmoid(beta_raw[token * heads + head]);
6930            let row_offset = (token * heads + head) * head_dim;
6931            let state_offset = head * head_dim * head_dim;
6932            for key_index in 0..head_dim {
6933                let decay = forget[token * qkv + head * head_dim + key_index].exp();
6934                let state_row = state_offset + key_index * head_dim;
6935                for value_index in 0..head_dim {
6936                    matrix[state_row + value_index] *= decay;
6937                }
6938            }
6939            let mut delta = vec![0.0; head_dim];
6940            for value_index in 0..head_dim {
6941                let mut memory = 0.0;
6942                for key_index in 0..head_dim {
6943                    memory += matrix[state_offset + key_index * head_dim + value_index]
6944                        * key[row_offset + key_index];
6945                }
6946                delta[value_index] = (value[row_offset + value_index] - memory) * beta;
6947            }
6948            for key_index in 0..head_dim {
6949                let state_row = state_offset + key_index * head_dim;
6950                for value_index in 0..head_dim {
6951                    matrix[state_row + value_index] +=
6952                        key[row_offset + key_index] * delta[value_index];
6953                }
6954            }
6955            for value_index in 0..head_dim {
6956                let mut attended = 0.0;
6957                for key_index in 0..head_dim {
6958                    attended += matrix[state_offset + key_index * head_dim + value_index]
6959                        * query[row_offset + key_index];
6960                }
6961                core[row_offset + value_index] = attended;
6962            }
6963        }
6964    }
6965
6966    // Output: sigmoid-gated fp32 RMSNorm over head_dim (o_norm uses the layer's
6967    // rms_norm_eps), gate = g_b(g_a(x)); then o_proj.
6968    let gate_down = linear(
6969        x,
6970        tensor(
6971            weights,
6972            &layer_id(layer, LayerTensor::KdaGateDown),
6973            &[head_dim, hidden],
6974        )?,
6975        tokens,
6976        hidden,
6977        head_dim,
6978    );
6979    let gate = linear(
6980        &gate_down,
6981        tensor(
6982            weights,
6983            &layer_id(layer, LayerTensor::KdaGateUp),
6984            &[qkv, head_dim],
6985        )?,
6986        tokens,
6987        head_dim,
6988        qkv,
6989    );
6990    let norm_weight = tensor(
6991        weights,
6992        &layer_id(layer, LayerTensor::KdaOutputNorm),
6993        &[head_dim],
6994    )?;
6995    let mut gated = rms_norm(&core, tokens * heads, head_dim, norm_weight, epsilon);
6996    for index in 0..gated.len() {
6997        gated[index] *= sigmoid(gate[index]);
6998    }
6999    let output = linear(
7000        &gated,
7001        tensor(
7002            weights,
7003            &layer_id(layer, LayerTensor::KdaOutput),
7004            &[hidden, qkv],
7005        )?,
7006        tokens,
7007        qkv,
7008        hidden,
7009    );
7010
7011    // Conv state stores the raw fused [q|k|v] pre-conv planes for the trailing
7012    // kernel-1 positions, mirroring the GDN layout.
7013    let pad = kernel - 1;
7014    let mut conv_state = vec![0.0; conv_width * pad];
7015    let planes = [&q_raw, &k_raw, &v_raw];
7016    for channel in 0..conv_width {
7017        let plane = channel / qkv;
7018        let plane_channel = channel % qkv;
7019        for index in 0..pad {
7020            let source = tokens as isize - pad as isize + index as isize;
7021            if source >= 0 {
7022                conv_state[channel * pad + index] =
7023                    planes[plane][source as usize * qkv + plane_channel];
7024            }
7025        }
7026    }
7027    Ok((
7028        output,
7029        ReferenceLayerState::Recurrent {
7030            conv: conv_state,
7031            matrix,
7032            value_heads: heads,
7033            key_head_dim: head_dim,
7034            value_head_dim: head_dim,
7035            conv_width,
7036        },
7037    ))
7038}
7039
7040#[allow(clippy::too_many_arguments)]
7041// allow: the parameter list mirrors the kernel/FFI/call contract; bundling into a struct is a refactor, not a lint fix
7042#[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
7043fn full_attention(
7044    layer: u32,
7045    plan: &memra_gguf::model_plan::FullAttentionPlan,
7046    window: Option<usize>,
7047    norm_epsilon: f32,
7048    weights: &ReferenceWeights,
7049    x: &[f32],
7050    tokens: usize,
7051    hidden: usize,
7052    // qwen4_exp QSA: indexer visibility overlay, `[tokens, tokens]` row-major (query,
7053    // source); attention runs dense under causal AND selection (SEMANTICS.md §QSA).
7054    selection: Option<&[bool]>,
7055) -> Result<(Vec<f32>, ReferenceLayerState), ReferenceError> {
7056    let query_heads = plan.query_heads as usize;
7057    let kv_heads = plan.kv_heads as usize;
7058    let key_dim = plan.key_head_dim as usize;
7059    let value_dim = plan.value_head_dim as usize;
7060    if query_heads == 0 || kv_heads == 0 || query_heads % kv_heads != 0 {
7061        return Err(ReferenceError::InvalidPlan {
7062            layer: Some(layer),
7063            reason: "query heads must be a positive multiple of KV heads",
7064        });
7065    }
7066    if selection.is_some_and(|selection| selection.len() != tokens * tokens) {
7067        return Err(ReferenceError::InvalidPlan {
7068            layer: Some(layer),
7069            reason: "attention selection mask does not match tokens x tokens",
7070        });
7071    }
7072    let fused = plan.output_gate == AttentionGateKind::FusedQ;
7073    let q_width = query_heads * key_dim;
7074    let q_projection_width = q_width * if fused { 2 } else { 1 };
7075    let k_width = kv_heads * key_dim;
7076    let v_width = kv_heads * value_dim;
7077    let q_weight = tensor(
7078        weights,
7079        &layer_id(layer, LayerTensor::Query),
7080        &[q_projection_width, hidden],
7081    )?;
7082    let k_weight = tensor(
7083        weights,
7084        &layer_id(layer, LayerTensor::Key),
7085        &[k_width, hidden],
7086    )?;
7087    let output_weight = tensor(
7088        weights,
7089        &layer_id(layer, LayerTensor::AttentionOutput),
7090        &[hidden, query_heads * value_dim],
7091    )?;
7092    let q_projected = linear(x, q_weight, tokens, hidden, q_projection_width);
7093    let mut query = vec![0.0; tokens * q_width];
7094    let mut fused_gate = None;
7095    if fused {
7096        let mut gate = vec![0.0; tokens * q_width];
7097        for token in 0..tokens {
7098            for head in 0..query_heads {
7099                let projected = token * q_projection_width + head * 2 * key_dim;
7100                let canonical = (token * query_heads + head) * key_dim;
7101                query[canonical..canonical + key_dim]
7102                    .copy_from_slice(&q_projected[projected..projected + key_dim]);
7103                gate[canonical..canonical + key_dim]
7104                    .copy_from_slice(&q_projected[projected + key_dim..projected + 2 * key_dim]);
7105            }
7106        }
7107        fused_gate = Some(gate);
7108    } else {
7109        query.copy_from_slice(&q_projected);
7110    }
7111    let mut key = linear(x, k_weight, tokens, hidden, k_width);
7112    let mut value = match plan.value_projection {
7113        ValueProjection::Separate => linear(
7114            x,
7115            tensor(
7116                weights,
7117                &layer_id(layer, LayerTensor::Value),
7118                &[v_width, hidden],
7119            )?,
7120            tokens,
7121            hidden,
7122            v_width,
7123        ),
7124        ValueProjection::ReuseKey => {
7125            if value_dim != key_dim {
7126                return Err(ReferenceError::InvalidPlan {
7127                    layer: Some(layer),
7128                    reason: "K-as-V requires equal key/value head widths",
7129                });
7130            }
7131            key.clone()
7132        }
7133    };
7134    apply_optional_head_norm(
7135        weights,
7136        layer_id(layer, LayerTensor::QueryNorm),
7137        &mut query,
7138        tokens * query_heads,
7139        key_dim,
7140        plan.qk_norm,
7141        norm_epsilon,
7142    )?;
7143    if plan.value_norm == ValueNorm::WeightlessRms {
7144        let ones = vec![1.0; value_dim];
7145        value = rms_norm(&value, tokens * kv_heads, value_dim, &ones, norm_epsilon);
7146    }
7147    apply_optional_head_norm(
7148        weights,
7149        layer_id(layer, LayerTensor::KeyNorm),
7150        &mut key,
7151        tokens * kv_heads,
7152        key_dim,
7153        plan.qk_norm,
7154        norm_epsilon,
7155    )?;
7156    let (rope_factors, rope_mscale) = rope_factor_values(&plan.rope, weights)?;
7157    apply_rope(
7158        &mut query,
7159        tokens,
7160        query_heads,
7161        key_dim,
7162        plan.rope.dimensions as usize,
7163        plan.rope.base,
7164        rope_factors.as_deref(),
7165        rope_mscale,
7166    );
7167    apply_rope(
7168        &mut key,
7169        tokens,
7170        kv_heads,
7171        key_dim,
7172        plan.rope.dimensions as usize,
7173        plan.rope.base,
7174        rope_factors.as_deref(),
7175        rope_mscale,
7176    );
7177
7178    let mut attended = vec![0.0; tokens * query_heads * value_dim];
7179    let scale = match plan.scale {
7180        AttentionScale::InverseSqrtKeyDim => 1.0 / (key_dim as f32).sqrt(),
7181        AttentionScale::Fixed(scale) => scale,
7182    };
7183    for token in 0..tokens {
7184        for head in 0..query_heads {
7185            let kv_head = head * kv_heads / query_heads;
7186            let first_source = window
7187                .map(|window| (token + 1).saturating_sub(window))
7188                .unwrap_or(0);
7189            let mut sources = Vec::with_capacity(token + 1 - first_source);
7190            let mut scores = Vec::with_capacity(token + 1 - first_source);
7191            for source in first_source..=token {
7192                if selection.is_some_and(|selection| !selection[token * tokens + source]) {
7193                    continue;
7194                }
7195                let mut score = 0.0;
7196                for dim in 0..key_dim {
7197                    score += query[(token * query_heads + head) * key_dim + dim]
7198                        * key[(source * kv_heads + kv_head) * key_dim + dim];
7199                }
7200                sources.push(source);
7201                scores.push(score * scale);
7202            }
7203            if scores.is_empty() {
7204                // The QSA tail rule guarantees every query keeps at least its own block's
7205                // incomplete tail; an empty row means a malformed selection mask.
7206                return Err(ReferenceError::InvalidPlan {
7207                    layer: Some(layer),
7208                    reason: "attention selection left a query with no visible source",
7209                });
7210            }
7211            softmax_in_place(&mut scores);
7212            for (index, probability) in scores.into_iter().enumerate() {
7213                let source = sources[index];
7214                for dim in 0..value_dim {
7215                    attended[(token * query_heads + head) * value_dim + dim] +=
7216                        probability * value[(source * kv_heads + kv_head) * value_dim + dim];
7217                }
7218            }
7219        }
7220    }
7221    if let Some(gate) = fused_gate {
7222        for token in 0..tokens {
7223            for head in 0..query_heads {
7224                for dim in 0..value_dim {
7225                    if dim >= key_dim {
7226                        return Err(ReferenceError::InvalidPlan {
7227                            layer: Some(layer),
7228                            reason: "fused attention gate requires value_dim <= key_dim",
7229                        });
7230                    }
7231                    attended[(token * query_heads + head) * value_dim + dim] *=
7232                        sigmoid(gate[(token * query_heads + head) * key_dim + dim]);
7233                }
7234            }
7235        }
7236    } else if plan.output_gate == AttentionGateKind::SeparateHead {
7237        let gate_weight = tensor(
7238            weights,
7239            &layer_id(layer, LayerTensor::AttentionGate),
7240            &[query_heads, hidden],
7241        )?;
7242        let gates = linear(x, gate_weight, tokens, hidden, query_heads);
7243        for token in 0..tokens {
7244            for head in 0..query_heads {
7245                let gate = sigmoid(gates[token * query_heads + head]);
7246                for dim in 0..value_dim {
7247                    attended[(token * query_heads + head) * value_dim + dim] *= gate;
7248                }
7249            }
7250        }
7251    }
7252    let state_start = window
7253        .map(|window| tokens.saturating_sub(window))
7254        .unwrap_or(0);
7255    let state_tokens = tokens - state_start;
7256    let state_key = key[state_start * k_width..].to_vec();
7257    let state_value = value[state_start * v_width..].to_vec();
7258    Ok((
7259        linear(
7260            &attended,
7261            output_weight,
7262            tokens,
7263            query_heads * value_dim,
7264            hidden,
7265        ),
7266        ReferenceLayerState::Kv {
7267            key: state_key,
7268            value: state_value,
7269            tokens: state_tokens,
7270            kv_heads,
7271            key_head_dim: key_dim,
7272            value_head_dim: value_dim,
7273            window,
7274        },
7275    ))
7276}
7277
7278fn dense_mlp(
7279    layer: u32,
7280    plan: &memra_gguf::model_plan::DenseMlpPlan,
7281    weights: &ReferenceWeights,
7282    x: &[f32],
7283    tokens: usize,
7284    hidden: usize,
7285) -> Result<Vec<f32>, ReferenceError> {
7286    let intermediate = plan.intermediate_size as usize;
7287    let gate = linear(
7288        x,
7289        tensor(
7290            weights,
7291            &layer_id(layer, LayerTensor::MlpGate),
7292            &[intermediate, hidden],
7293        )?,
7294        tokens,
7295        hidden,
7296        intermediate,
7297    );
7298    let up = linear(
7299        x,
7300        tensor(
7301            weights,
7302            &layer_id(layer, LayerTensor::MlpUp),
7303            &[intermediate, hidden],
7304        )?,
7305        tokens,
7306        hidden,
7307        intermediate,
7308    );
7309    let mut activated = vec![0.0; gate.len()];
7310    for index in 0..gate.len() {
7311        activated[index] = activate_pair(&plan.activation, gate[index], up[index], layer)?;
7312    }
7313    Ok(linear(
7314        &activated,
7315        tensor(
7316            weights,
7317            &layer_id(layer, LayerTensor::MlpDown),
7318            &[hidden, intermediate],
7319        )?,
7320        tokens,
7321        intermediate,
7322        hidden,
7323    ))
7324}
7325
7326#[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
7327fn moe_mlp(
7328    layer: u32,
7329    plan: &memra_gguf::model_plan::MoeMlpPlan,
7330    weights: &ReferenceWeights,
7331    x: &[f32],
7332    token_ids: &[u32],
7333    tokens: usize,
7334    hidden: usize,
7335    vocab: usize,
7336) -> Result<Vec<f32>, ReferenceError> {
7337    let experts = plan.expert_count as usize;
7338    let selected = plan.experts_per_token as usize;
7339    let intermediate = plan.expert_intermediate_size as usize;
7340    if selected == 0 || selected > experts {
7341        return Err(ReferenceError::InvalidPlan {
7342            layer: Some(layer),
7343            reason: "MoE top-k must be in 1..=expert_count",
7344        });
7345    }
7346    let router = tensor(
7347        weights,
7348        &layer_id(layer, LayerTensor::MoeRouter),
7349        &[experts, hidden],
7350    )?;
7351    let logits = linear(x, router, tokens, hidden, experts);
7352    let bias = if router_has_selection_bias(&plan.router) {
7353        Some(tensor(
7354            weights,
7355            &layer_id(layer, LayerTensor::MoeRouterBias),
7356            &[experts],
7357        )?)
7358    } else {
7359        None
7360    };
7361    let token_to_expert = if matches!(
7362        plan.router,
7363        memra_gguf::model_plan::RouterPlan::TokenIdHash { .. }
7364    ) {
7365        Some(tensor(
7366            weights,
7367            &layer_id(layer, LayerTensor::MoeTokenToExpert),
7368            &[vocab, selected],
7369        )?)
7370    } else {
7371        None
7372    };
7373    let gate_bank = tensor(
7374        weights,
7375        &layer_id(layer, LayerTensor::MoeExpertGateBank),
7376        &[experts, intermediate, hidden],
7377    )?;
7378    let up_bank = tensor(
7379        weights,
7380        &layer_id(layer, LayerTensor::MoeExpertUpBank),
7381        &[experts, intermediate, hidden],
7382    )?;
7383    let down_bank = tensor(
7384        weights,
7385        &layer_id(layer, LayerTensor::MoeExpertDownBank),
7386        &[experts, hidden, intermediate],
7387    )?;
7388    let mut output = vec![0.0; tokens * hidden];
7389    for token in 0..tokens {
7390        let forced_routes = token_to_expert
7391            .map(|table| {
7392                let token_id = token_ids[token] as usize;
7393                &table[token_id * selected..(token_id + 1) * selected]
7394            })
7395            .map(|row| {
7396                row.iter()
7397                    .map(|&value| {
7398                        if !value.is_finite()
7399                            || value < 0.0
7400                            || value.fract() != 0.0
7401                            || value as usize >= experts
7402                        {
7403                            return Err(ReferenceError::InvalidPlan {
7404                                layer: Some(layer),
7405                                reason: "token-id expert table contains an invalid expert id",
7406                            });
7407                        }
7408                        Ok(value as usize)
7409                    })
7410                    .collect::<Result<Vec<_>, _>>()
7411            })
7412            .transpose()?;
7413        let routes = route_experts(
7414            &plan.router,
7415            &logits[token * experts..(token + 1) * experts],
7416            bias,
7417            selected,
7418            forced_routes.as_deref(),
7419            layer,
7420        )?;
7421        if crate::hidden_trace::enabled() && token + 1 == tokens {
7422            crate::hidden_trace::emit_last_row(
7423                "router",
7424                layer as i64,
7425                1,
7426                experts,
7427                &logits[token * experts..(token + 1) * experts],
7428            );
7429            let mut route = Vec::with_capacity(routes.len() * 2);
7430            for (expert, weight) in &routes {
7431                route.push(*expert as f32);
7432                route.push(*weight);
7433            }
7434            crate::hidden_trace::emit_last_row("route", layer as i64, 1, route.len(), &route);
7435        }
7436        let input = &x[token * hidden..(token + 1) * hidden];
7437        for (expert, route_weight) in routes {
7438            let gate_offset = expert * intermediate * hidden;
7439            let down_offset = expert * hidden * intermediate;
7440            let mut activated = vec![0.0; intermediate];
7441            for row in 0..intermediate {
7442                let mut gate = 0.0;
7443                let mut up = 0.0;
7444                for column in 0..hidden {
7445                    gate += input[column] * gate_bank[gate_offset + row * hidden + column];
7446                    up += input[column] * up_bank[gate_offset + row * hidden + column];
7447                }
7448                activated[row] = activate_pair(&plan.activation, gate, up, layer)?;
7449            }
7450            for row in 0..hidden {
7451                let mut value = 0.0;
7452                for column in 0..intermediate {
7453                    value +=
7454                        activated[column] * down_bank[down_offset + row * intermediate + column];
7455                }
7456                output[token * hidden + row] += route_weight * value;
7457            }
7458        }
7459    }
7460
7461    if crate::hidden_trace::enabled() {
7462        crate::hidden_trace::emit_last_row("routed", layer as i64, tokens, hidden, &output);
7463    }
7464
7465    if let Some(shared) = plan.shared.as_ref() {
7466        let intermediate = shared.intermediate_size as usize;
7467        let gate = linear(
7468            x,
7469            tensor(
7470                weights,
7471                &layer_id(layer, LayerTensor::SharedMlpGate),
7472                &[intermediate, hidden],
7473            )?,
7474            tokens,
7475            hidden,
7476            intermediate,
7477        );
7478        let up = linear(
7479            x,
7480            tensor(
7481                weights,
7482                &layer_id(layer, LayerTensor::SharedMlpUp),
7483                &[intermediate, hidden],
7484            )?,
7485            tokens,
7486            hidden,
7487            intermediate,
7488        );
7489        let mut activated = vec![0.0; gate.len()];
7490        for index in 0..gate.len() {
7491            activated[index] = activate_pair(&plan.activation, gate[index], up[index], layer)?;
7492        }
7493        let mut shared_output = linear(
7494            &activated,
7495            tensor(
7496                weights,
7497                &layer_id(layer, LayerTensor::SharedMlpDown),
7498                &[hidden, intermediate],
7499            )?,
7500            tokens,
7501            intermediate,
7502            hidden,
7503        );
7504        if shared.gated {
7505            let gate_weight = tensor(
7506                weights,
7507                &layer_id(layer, LayerTensor::SharedMlpInputGate),
7508                &[hidden],
7509            )?;
7510            for token in 0..tokens {
7511                let mut gate = 0.0;
7512                for column in 0..hidden {
7513                    gate += x[token * hidden + column] * gate_weight[column];
7514                }
7515                let gate = sigmoid(gate);
7516                for column in 0..hidden {
7517                    shared_output[token * hidden + column] *= gate;
7518                }
7519            }
7520        }
7521        add_in_place(&mut output, &shared_output);
7522    }
7523    Ok(output)
7524}
7525
7526fn route_experts(
7527    router: &memra_gguf::model_plan::RouterPlan,
7528    logits: &[f32],
7529    bias: Option<&[f32]>,
7530    selected: usize,
7531    forced_indices: Option<&[usize]>,
7532    layer: u32,
7533) -> Result<Vec<(usize, f32)>, ReferenceError> {
7534    use memra_gguf::model_plan::{RouterPlan, RouterScorePlan};
7535
7536    let mut weights = match router {
7537        RouterPlan::Softmax => {
7538            let mut probabilities = logits.to_vec();
7539            softmax_in_place(&mut probabilities);
7540            probabilities
7541        }
7542        RouterPlan::Sigmoid { .. } => logits.iter().map(|&value| sigmoid(value)).collect(),
7543        RouterPlan::SqrtSoftplus { .. } => {
7544            logits.iter().map(|&value| softplus(value).sqrt()).collect()
7545        }
7546        RouterPlan::TokenIdHash { score, .. } => match score {
7547            RouterScorePlan::Softmax => {
7548                let mut probabilities = logits.to_vec();
7549                softmax_in_place(&mut probabilities);
7550                probabilities
7551            }
7552            RouterScorePlan::Sigmoid => logits.iter().map(|&value| sigmoid(value)).collect(),
7553            RouterScorePlan::SqrtSoftplus => {
7554                logits.iter().map(|&value| softplus(value).sqrt()).collect()
7555            }
7556        },
7557    };
7558    let selection_scores: Vec<f32> = weights
7559        .iter()
7560        .enumerate()
7561        .map(|(index, &weight)| weight + bias.map_or(0.0, |bias| bias[index]))
7562        .collect();
7563    let indices = if let RouterPlan::TokenIdHash { .. } = router {
7564        let Some(forced) = forced_indices else {
7565            return Err(ReferenceError::InvalidPlan {
7566                layer: Some(layer),
7567                reason: "token-id hash router requires a token-to-expert row",
7568            });
7569        };
7570        if forced.len() != selected {
7571            return Err(ReferenceError::InvalidPlan {
7572                layer: Some(layer),
7573                reason: "token-id expert row width does not match MoE top-k",
7574            });
7575        }
7576        let mut seen = std::collections::BTreeSet::new();
7577        for &index in forced {
7578            if index >= logits.len() || !seen.insert(index) {
7579                return Err(ReferenceError::InvalidPlan {
7580                    layer: Some(layer),
7581                    reason: "token-id expert row contains an out-of-range or duplicate expert",
7582                });
7583            }
7584        }
7585        forced.to_vec()
7586    } else {
7587        if forced_indices.is_some() {
7588            return Err(ReferenceError::InvalidPlan {
7589                layer: Some(layer),
7590                reason: "score-selected router received forced expert indices",
7591            });
7592        }
7593        let mut indices: Vec<usize> = (0..logits.len()).collect();
7594        indices.sort_by(|&left, &right| {
7595            selection_scores[right]
7596                .total_cmp(&selection_scores[left])
7597                .then(left.cmp(&right))
7598        });
7599        indices.truncate(selected);
7600        indices
7601    };
7602    let (normalize, scaling) = match router {
7603        RouterPlan::Softmax => (true, 1.0),
7604        RouterPlan::Sigmoid {
7605            normalize_selected,
7606            scaling_factor,
7607            ..
7608        }
7609        | RouterPlan::SqrtSoftplus {
7610            normalize_selected,
7611            scaling_factor,
7612            ..
7613        } => (*normalize_selected, *scaling_factor),
7614        RouterPlan::TokenIdHash {
7615            normalize_selected,
7616            scaling_factor,
7617            ..
7618        } => (*normalize_selected, *scaling_factor),
7619    };
7620    if normalize {
7621        let denominator = indices
7622            .iter()
7623            .map(|&index| weights[index])
7624            .sum::<f32>()
7625            .max(if matches!(router, RouterPlan::Softmax) {
7626                6.103_515_6e-5
7627            } else {
7628                1e-20
7629            });
7630        for weight in &mut weights {
7631            *weight = *weight / denominator * scaling;
7632        }
7633    } else {
7634        for weight in &mut weights {
7635            *weight *= scaling;
7636        }
7637    }
7638    Ok(indices
7639        .into_iter()
7640        .map(|index| (index, weights[index]))
7641        .collect())
7642}
7643
7644fn router_has_selection_bias(router: &memra_gguf::model_plan::RouterPlan) -> bool {
7645    matches!(
7646        router,
7647        memra_gguf::model_plan::RouterPlan::Sigmoid {
7648            selection_bias: true,
7649            ..
7650        } | memra_gguf::model_plan::RouterPlan::SqrtSoftplus {
7651            selection_bias: true,
7652            ..
7653        }
7654    )
7655}
7656
7657fn activate_pair(
7658    activation: &ActivationPlan,
7659    gate: f32,
7660    up: f32,
7661    layer: u32,
7662) -> Result<f32, ReferenceError> {
7663    Ok(match activation {
7664        ActivationPlan::Silu => silu(gate) * up,
7665        ActivationPlan::GeluTanh => gelu_tanh(gate) * up,
7666        ActivationPlan::SwiGluOai { alpha, limit } => {
7667            (gate * sigmoid(*alpha * gate)).min(*limit) * up.clamp(-*limit, *limit)
7668        }
7669        ActivationPlan::SwiGluClamped { limit } => {
7670            silu(gate).min(*limit) * up.clamp(-*limit, *limit)
7671        }
7672        // glm5_next: the gate clamp is PRE-silu and one-sided (no lower bound).
7673        ActivationPlan::SwiGluPreClamped { limit } => {
7674            silu(gate.min(*limit)) * up.clamp(-*limit, *limit)
7675        }
7676        ActivationPlan::Named(_) => {
7677            return Err(ReferenceError::UnsupportedOperation {
7678                layer: Some(layer),
7679                operation: "named MLP activation",
7680            });
7681        }
7682    })
7683}
7684
7685fn tensor<'a>(
7686    weights: &'a ReferenceWeights,
7687    id: &TensorId,
7688    expected: &[usize],
7689) -> Result<&'a [f32], ReferenceError> {
7690    let tensor = weights
7691        .get(id)
7692        .ok_or_else(|| ReferenceError::MissingTensor(id.clone()))?;
7693    tensor_checked(id, tensor, expected)
7694}
7695
7696fn tensor_checked<'a>(
7697    id: &TensorId,
7698    tensor: &'a ReferenceTensor,
7699    expected: &[usize],
7700) -> Result<&'a [f32], ReferenceError> {
7701    if tensor.shape != expected {
7702        return Err(ReferenceError::TensorShape {
7703            id: Some(id.clone()),
7704            expected: expected.to_vec(),
7705            actual_elements: tensor.data.len(),
7706        });
7707    }
7708    Ok(&tensor.data)
7709}
7710
7711fn layer_id(layer: u32, tensor: LayerTensor) -> TensorId {
7712    TensorId::Layer {
7713        index: layer,
7714        tensor,
7715    }
7716}
7717
7718fn linear(x: &[f32], weight: &[f32], rows: usize, input: usize, output: usize) -> Vec<f32> {
7719    let mut result = vec![0.0; rows * output];
7720    for row in 0..rows {
7721        for out in 0..output {
7722            let mut sum = 0.0;
7723            for inner in 0..input {
7724                sum += x[row * input + inner] * weight[out * input + inner];
7725            }
7726            result[row * output + out] = sum;
7727        }
7728    }
7729    result
7730}
7731
7732fn rms_norm(x: &[f32], rows: usize, width: usize, weight: &[f32], epsilon: f32) -> Vec<f32> {
7733    let mut result = vec![0.0; x.len()];
7734    for row in 0..rows {
7735        let input = &x[row * width..(row + 1) * width];
7736        let mean_square = input.iter().map(|value| value * value).sum::<f32>() / width as f32;
7737        let inverse = 1.0 / (mean_square + epsilon).sqrt();
7738        for index in 0..width {
7739            result[row * width + index] = input[index] * inverse * weight[index];
7740        }
7741    }
7742    result
7743}
7744
7745/// LayerNorm WITH bias (indexer k_norm). The epsilon is nn.LayerNorm's default and
7746/// does NOT track rms_norm_eps.
7747fn layer_norm(x: &[f32], rows: usize, width: usize, weight: &[f32], bias: &[f32]) -> Vec<f32> {
7748    const EPSILON: f32 = 1e-5;
7749    let mut result = vec![0.0; x.len()];
7750    for row in 0..rows {
7751        let input = &x[row * width..(row + 1) * width];
7752        let mean = input.iter().sum::<f32>() / width as f32;
7753        let variance = input
7754            .iter()
7755            .map(|value| (value - mean) * (value - mean))
7756            .sum::<f32>()
7757            / width as f32;
7758        let inverse = 1.0 / (variance + EPSILON).sqrt();
7759        for index in 0..width {
7760            result[row * width + index] =
7761                (input[index] - mean) * inverse * weight[index] + bias[index];
7762        }
7763    }
7764    result
7765}
7766
7767fn l2_normalize_rows(values: &mut [f32], rows: usize, width: usize, epsilon: f32) {
7768    for row in 0..rows {
7769        let offset = row * width;
7770        let sum = values[offset..offset + width]
7771            .iter()
7772            .map(|value| value * value)
7773            .sum::<f32>();
7774        let inverse = 1.0 / (sum + epsilon).sqrt();
7775        for value in &mut values[offset..offset + width] {
7776            *value *= inverse;
7777        }
7778    }
7779}
7780
7781fn apply_optional_head_norm(
7782    weights: &ReferenceWeights,
7783    id: TensorId,
7784    values: &mut [f32],
7785    rows: usize,
7786    width: usize,
7787    presence: memra_gguf::model_plan::TensorPresence,
7788    epsilon: f32,
7789) -> Result<(), ReferenceError> {
7790    let Some(weight) = weights.get(&id) else {
7791        return if presence == memra_gguf::model_plan::TensorPresence::Required {
7792            Err(ReferenceError::MissingTensor(id))
7793        } else {
7794            Ok(())
7795        };
7796    };
7797    let normalized = rms_norm(
7798        values,
7799        rows,
7800        width,
7801        tensor_checked(&id, weight, &[width])?,
7802        epsilon,
7803    );
7804    values.copy_from_slice(&normalized);
7805    Ok(())
7806}
7807
7808/// Per-dim frequency divisors + the cos/sin attention scale (YaRN mscale; 1.0 for every
7809/// other factor kind — an exact multiplicative identity).
7810fn rope_factor_values(
7811    plan: &memra_gguf::model_plan::RopePlan,
7812    weights: &ReferenceWeights,
7813) -> Result<(Option<Vec<f32>>, f32), ReferenceError> {
7814    use memra_gguf::model_plan::RopeFactors;
7815
7816    let width = plan.dimensions as usize / 2;
7817    Ok(match plan.factors {
7818        RopeFactors::None => (None, 1.0),
7819        RopeFactors::PartialRotary { factor } => {
7820            let keep = (width as f32 * factor.clamp(0.0, 1.0)).round() as usize;
7821            (
7822                Some(
7823                    (0..width)
7824                        .map(|index| if index < keep { 1.0 } else { 1.0e30 })
7825                        .collect(),
7826                ),
7827                1.0,
7828            )
7829        }
7830        RopeFactors::Checkpoint => {
7831            let tensor = weights
7832                .get(&TensorId::RopeFactors)
7833                .ok_or(ReferenceError::MissingTensor(TensorId::RopeFactors))?;
7834            if tensor.shape.len() != 1 || tensor.data.len() < width {
7835                return Err(ReferenceError::TensorShape {
7836                    id: Some(TensorId::RopeFactors),
7837                    expected: vec![width],
7838                    actual_elements: tensor.data.len(),
7839                });
7840            }
7841            (Some(tensor.data[..width].to_vec()), 1.0)
7842        }
7843        // YaRN on full attention (qwen4_exp long-context lane): the transformers-twin
7844        // frequency divisors + the derived attention factor on cos/sin. The divisor table
7845        // shares the Checkpoint-factors convention, so every consumer below (QSA q/k AND
7846        // the indexer's q/pooled-k rope) rides the same path.
7847        RopeFactors::Yarn {
7848            factor,
7849            original_context,
7850            beta_fast,
7851            beta_slow,
7852        } => (
7853            Some(memra_gguf::model_plan::yarn_frequency_divisors(
7854                plan.dimensions,
7855                plan.base,
7856                factor,
7857                original_context,
7858                beta_fast,
7859                beta_slow,
7860            )),
7861            memra_gguf::model_plan::yarn_attention_factor(factor),
7862        ),
7863    })
7864}
7865
7866#[allow(clippy::too_many_arguments)]
7867fn apply_rope(
7868    values: &mut [f32],
7869    tokens: usize,
7870    heads: usize,
7871    head_dim: usize,
7872    dimensions: usize,
7873    base: f32,
7874    factors: Option<&[f32]>,
7875    mscale: f32,
7876) {
7877    for token in 0..tokens {
7878        apply_rope_at_position(
7879            &mut values[token * heads * head_dim..(token + 1) * heads * head_dim],
7880            heads,
7881            head_dim,
7882            dimensions,
7883            base,
7884            factors,
7885            mscale,
7886            token,
7887        );
7888    }
7889}
7890
7891/// One row of NeoX split-half rope at an EXPLICIT position — the QSA indexer rotates
7892/// pooled block keys at the block-start position, not their row index. `mscale` is the
7893/// YaRN attention factor on cos/sin (transformers `attention_scaling`; 1.0 elsewhere —
7894/// an exact multiplicative identity, so the non-yarn arms are byte-unchanged).
7895#[allow(clippy::too_many_arguments)]
7896fn apply_rope_at_position(
7897    values: &mut [f32],
7898    heads: usize,
7899    head_dim: usize,
7900    dimensions: usize,
7901    base: f32,
7902    factors: Option<&[f32]>,
7903    mscale: f32,
7904    position: usize,
7905) {
7906    let dimensions = dimensions.min(head_dim) / 2 * 2;
7907    let half = dimensions / 2;
7908    for head in 0..heads {
7909        let offset = head * head_dim;
7910        for index in 0..half {
7911            let factor = factors.map_or(1.0, |factors| factors[index]);
7912            let frequency = base.powf(-2.0 * index as f32 / dimensions as f32) / factor;
7913            let angle = position as f32 * frequency;
7914            let (sin, cos) = angle.sin_cos();
7915            let (sin, cos) = (sin * mscale, cos * mscale);
7916            let first = values[offset + index];
7917            let second = values[offset + index + half];
7918            values[offset + index] = first * cos - second * sin;
7919            values[offset + index + half] = first * sin + second * cos;
7920        }
7921    }
7922}
7923
7924fn softmax_in_place(values: &mut [f32]) {
7925    let max = values.iter().copied().fold(f32::NEG_INFINITY, f32::max);
7926    let mut sum = 0.0;
7927    for value in values.iter_mut() {
7928        *value = (*value - max).exp();
7929        sum += *value;
7930    }
7931    for value in values {
7932        *value /= sum;
7933    }
7934}
7935
7936fn add_in_place(target: &mut [f32], addend: &[f32]) {
7937    for (target, addend) in target.iter_mut().zip(addend) {
7938        *target += addend;
7939    }
7940}
7941
7942fn sigmoid(value: f32) -> f32 {
7943    1.0 / (1.0 + (-value).exp())
7944}
7945
7946fn silu(value: f32) -> f32 {
7947    value * sigmoid(value)
7948}
7949
7950fn softplus(value: f32) -> f32 {
7951    if value > 20.0 {
7952        value
7953    } else {
7954        value.exp().ln_1p()
7955    }
7956}
7957
7958fn gelu_tanh(value: f32) -> f32 {
7959    0.5 * value * (1.0 + (0.797_884_6 * (value + 0.044_715 * value * value * value)).tanh())
7960}
7961
7962/// Exact-erf GELU (torch `nn.GELU()` default, used by the glm5_next vision merger — NOT
7963/// the tanh approximation). erf via Abramowitz & Stegun 7.1.26 in f64 (max abs error
7964/// 1.5e-7, below f32 resolution at these magnitudes).
7965fn gelu_erf(value: f32) -> f32 {
7966    let x = value as f64 / std::f64::consts::SQRT_2;
7967    let sign = if x < 0.0 { -1.0 } else { 1.0 };
7968    let x = x.abs();
7969    let t = 1.0 / (1.0 + 0.327_591_1 * x);
7970    let poly = t
7971        * (0.254_829_592
7972            + t * (-0.284_496_736
7973                + t * (1.421_413_741 + t * (-1.453_152_027 + t * 1.061_405_429))));
7974    let erf = sign * (1.0 - poly * (-x * x).exp());
7975    (0.5 * value as f64 * (1.0 + erf)) as f32
7976}
7977
7978#[cfg(test)]
7979mod tests {
7980    use super::*;
7981    use memra_gguf::config::{HfConfig, ModelConfig};
7982
7983    fn weight(shape: &[usize], data: &[f32]) -> ReferenceTensor {
7984        ReferenceTensor::new(shape.to_vec(), data.to_vec()).unwrap()
7985    }
7986
7987    #[test]
7988    fn one_token_dense_plan_matches_hand_derived_logits_and_emits_kv_state() {
7989        let config = ModelConfig::from_hf(&HfConfig::parse(
7990            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":2,
7991            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":2,
7992            "intermediate_size":2,"vocab_size":3,"max_position_embeddings":8,
7993            "rms_norm_eps":0.000001}"#,
7994        ));
7995        let plan = ModelPlan::compile(&config).unwrap();
7996        let identity = [1.0, 0.0, 0.0, 1.0];
7997        let zero = [0.0; 4];
7998        let mut weights = ReferenceWeights::new();
7999        weights.insert(
8000            TensorId::TokenEmbedding,
8001            weight(&[3, 2], &[1.0, 0.0, 0.0, 1.0, -1.0, 0.0]),
8002        );
8003        weights.insert(TensorId::OutputNorm, weight(&[2], &[1.0, 1.0]));
8004        for tensor in [LayerTensor::PreAttentionNorm, LayerTensor::PreMlpNorm] {
8005            weights.insert(layer_id(0, tensor), weight(&[2], &[1.0, 1.0]));
8006        }
8007        for tensor in [
8008            LayerTensor::Query,
8009            LayerTensor::Key,
8010            LayerTensor::Value,
8011            LayerTensor::AttentionOutput,
8012        ] {
8013            weights.insert(layer_id(0, tensor), weight(&[2, 2], &identity));
8014        }
8015        for tensor in [
8016            LayerTensor::MlpGate,
8017            LayerTensor::MlpUp,
8018            LayerTensor::MlpDown,
8019        ] {
8020            weights.insert(layer_id(0, tensor), weight(&[2, 2], &zero));
8021        }
8022
8023        let output = execute(&plan, &weights, &[0]).unwrap();
8024        let root_two = 2.0f32.sqrt();
8025        assert_eq!((output.tokens, output.vocab), (1, 3));
8026        assert!((output.logits[0] - root_two).abs() < 2e-5);
8027        assert!(output.logits[1].abs() < 2e-5);
8028        assert!((output.logits[2] + root_two).abs() < 2e-5);
8029        let ReferenceLayerState::Kv {
8030            tokens, key, value, ..
8031        } = &output.state.layers[0]
8032        else {
8033            panic!("expected KV state");
8034        };
8035        assert_eq!(*tokens, 1);
8036        assert_eq!(key.len(), 2);
8037        assert_eq!(value.len(), 2);
8038    }
8039
8040    #[test]
8041    fn hyperconnections_execute_stream_state_and_head_collapse() {
8042        let config = ModelConfig::from_hf(&HfConfig::parse(
8043            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":2,
8044            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":2,
8045            "intermediate_size":2,"vocab_size":3,"max_position_embeddings":8}"#,
8046        ));
8047        let mut plan = ModelPlan::compile(&config).unwrap();
8048        plan.layers[0].residual = ResidualTopology::HyperConnections {
8049            streams: 2,
8050            epsilon: 1e-6,
8051            sinkhorn_iterations: 2,
8052            collapse: HcCollapse::GatedHead,
8053        };
8054        let fixture = deterministic_fixture(&plan).unwrap();
8055        assert_eq!(
8056            fixture.weights[&TensorId::HyperHeadFunction].shape,
8057            vec![2, 4]
8058        );
8059        assert_eq!(
8060            fixture.weights[&layer_id(0, LayerTensor::HyperAttentionFunction)].shape,
8061            vec![8, 4]
8062        );
8063        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8064        assert!(output.logits.iter().all(|value| value.is_finite()));
8065        assert!(matches!(
8066            output.state.layers[0],
8067            ReferenceLayerState::Kv { .. }
8068        ));
8069    }
8070
8071    #[test]
8072    fn generated_tiny_fixture_is_deterministic_and_executable() {
8073        let config = ModelConfig::from_hf(&HfConfig::parse(
8074            r#"{"model_type":"qwen3","num_hidden_layers":2,"hidden_size":8,
8075            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8076            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32}"#,
8077        ));
8078        let plan = ModelPlan::compile(&config).unwrap();
8079        let first = deterministic_fixture(&plan).unwrap();
8080        let second = deterministic_fixture(&plan).unwrap();
8081        assert_eq!(first, second);
8082        let output = execute(&plan, &first.weights, &first.token_ids).unwrap();
8083        assert_eq!(output.logits.len(), first.token_ids.len() * 32);
8084        assert!(output.logits.iter().all(|value| value.is_finite()));
8085    }
8086
8087    #[test]
8088    fn qwen35_fixture_executes_mixed_gdn_and_full_attention_state() {
8089        let config = ModelConfig::from_hf(&HfConfig::parse(
8090            r#"{"model_type":"qwen3_5","num_hidden_layers":4,"hidden_size":8,
8091            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8092            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
8093            "rms_norm_eps":0.000001,"full_attention_interval":2,
8094            "linear_conv_kernel_dim":3,"linear_key_head_dim":4,
8095            "linear_value_head_dim":4,"linear_num_key_heads":1,
8096            "linear_num_value_heads":2}"#,
8097        ));
8098        let plan = ModelPlan::compile(&config).unwrap();
8099        let fixture = deterministic_fixture(&plan).unwrap();
8100        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8101        assert_eq!(output.state.layers.len(), 4);
8102        assert!(matches!(
8103            output.state.layers[0],
8104            ReferenceLayerState::Recurrent { .. }
8105        ));
8106        assert!(matches!(
8107            output.state.layers[1],
8108            ReferenceLayerState::Kv { .. }
8109        ));
8110        assert!(matches!(
8111            output.state.layers[2],
8112            ReferenceLayerState::Recurrent { .. }
8113        ));
8114        assert!(matches!(
8115            output.state.layers[3],
8116            ReferenceLayerState::Kv { .. }
8117        ));
8118        assert!(output.logits.iter().all(|value| value.is_finite()));
8119        assert_eq!(
8120            output.logits[..8]
8121                .iter()
8122                .map(|value| value.to_bits())
8123                .collect::<Vec<_>>(),
8124            vec![
8125                3_182_242_076,
8126                1_053_299_392,
8127                3_199_800_546,
8128                3_198_737_445,
8129                3_180_184_136,
8130                3_187_768_631,
8131                1_057_556_100,
8132                1_035_812_924,
8133            ]
8134        );
8135    }
8136
8137    #[test]
8138    fn router_laws_pin_stable_ties_and_selection_only_bias() {
8139        use memra_gguf::model_plan::{RouterPlan, RouterScorePlan};
8140
8141        assert_eq!(
8142            route_experts(&RouterPlan::Softmax, &[0.0, 0.0, 0.0], None, 2, None, 0,).unwrap(),
8143            vec![(0, 0.5), (1, 0.5)]
8144        );
8145        assert_eq!(
8146            route_experts(
8147                &RouterPlan::Sigmoid {
8148                    normalize_selected: true,
8149                    scaling_factor: 2.0,
8150                    selection_bias: true,
8151                },
8152                &[0.0, 0.0],
8153                Some(&[-1.0, 1.0]),
8154                1,
8155                None,
8156                0,
8157            )
8158            .unwrap(),
8159            vec![(1, 2.0)]
8160        );
8161        assert_eq!(
8162            route_experts(
8163                &RouterPlan::TokenIdHash {
8164                    score: RouterScorePlan::SqrtSoftplus,
8165                    normalize_selected: true,
8166                    scaling_factor: 1.5,
8167                },
8168                &[0.0, 0.0, 0.0],
8169                None,
8170                2,
8171                Some(&[2, 0]),
8172                0,
8173            )
8174            .unwrap(),
8175            vec![(2, 0.75), (0, 0.75)]
8176        );
8177        assert!(matches!(
8178            route_experts(
8179                &RouterPlan::TokenIdHash {
8180                    score: RouterScorePlan::SqrtSoftplus,
8181                    normalize_selected: true,
8182                    scaling_factor: 1.5,
8183                },
8184                &[0.0, 0.0, 0.0],
8185                None,
8186                2,
8187                Some(&[1, 1]),
8188                0,
8189            ),
8190            Err(ReferenceError::InvalidPlan {
8191                reason: "token-id expert row contains an out-of-range or duplicate expert",
8192                ..
8193            })
8194        ));
8195    }
8196
8197    #[test]
8198    fn token_hash_moe_fixture_executes_from_semantic_token_table() {
8199        use memra_gguf::model_plan::{RouterPlan, RouterScorePlan};
8200
8201        let config = ModelConfig::from_hf(&HfConfig::parse(
8202            r#"{"model_type":"qwen3_moe","num_hidden_layers":1,"hidden_size":8,
8203            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8204            "intermediate_size":16,"vocab_size":16,"max_position_embeddings":32,
8205            "num_experts":4,"num_experts_per_tok":2,"moe_intermediate_size":8}"#,
8206        ));
8207        let mut plan = ModelPlan::compile(&config).unwrap();
8208        let MlpPlan::Moe(moe) = &mut plan.layers[0].mlp else {
8209            unreachable!()
8210        };
8211        moe.router = RouterPlan::TokenIdHash {
8212            score: RouterScorePlan::SqrtSoftplus,
8213            normalize_selected: true,
8214            scaling_factor: 1.5,
8215        };
8216        let fixture = deterministic_fixture(&plan).unwrap();
8217        let table_id = layer_id(0, LayerTensor::MoeTokenToExpert);
8218        assert_eq!(fixture.weights[&table_id].shape, vec![16, 2]);
8219        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8220        assert!(output.logits.iter().all(|value| value.is_finite()));
8221
8222        let mut alternate = fixture.weights.clone();
8223        alternate.get_mut(&table_id).unwrap().data.fill(3.0);
8224        for row in alternate
8225            .get_mut(&table_id)
8226            .unwrap()
8227            .data
8228            .chunks_exact_mut(2)
8229        {
8230            row[1] = 2.0;
8231        }
8232        let alternate = execute(&plan, &alternate, &fixture.token_ids).unwrap();
8233        assert_ne!(output.logits, alternate.logits);
8234    }
8235
8236    #[test]
8237    fn qwen3_moe_fixture_executes_routed_and_shared_branches() {
8238        let config = ModelConfig::from_hf(&HfConfig::parse(
8239            r#"{"model_type":"qwen3_moe","num_hidden_layers":2,"hidden_size":8,
8240            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8241            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
8242            "num_experts":4,"num_experts_per_tok":2,"moe_intermediate_size":8,
8243            "shared_expert_intermediate_size":8}"#,
8244        ));
8245        let plan = ModelPlan::compile(&config).unwrap();
8246        let fixture = deterministic_fixture(&plan).unwrap();
8247        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8248        assert!(output.logits.iter().all(|value| value.is_finite()));
8249        assert_eq!(
8250            output.logits[..8]
8251                .iter()
8252                .map(|value| value.to_bits())
8253                .collect::<Vec<_>>(),
8254            vec![
8255                3_205_834_204,
8256                1_034_800_117,
8257                1_053_917_366,
8258                3_190_866_844,
8259                984_171_488,
8260                3_182_514_784,
8261                3_154_736_064,
8262                3_175_624_690,
8263            ]
8264        );
8265    }
8266
8267    #[test]
8268    fn sliding_window_limits_attention_and_trims_reference_state() {
8269        let config = ModelConfig::from_hf(&HfConfig::parse(
8270            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":8,
8271            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8272            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32}"#,
8273        ));
8274        let mut plan = ModelPlan::compile(&config).unwrap();
8275        let AttentionPlan::Full(attention) = plan.layers[0].attention.clone() else {
8276            unreachable!()
8277        };
8278        plan.layers[0].attention = AttentionPlan::SlidingWindow {
8279            attention,
8280            window: 2,
8281        };
8282        let fixture = deterministic_fixture(&plan).unwrap();
8283        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8284        let ReferenceLayerState::Kv { tokens, window, .. } = output.state.layers[0] else {
8285            panic!("expected sliding KV state");
8286        };
8287        assert_eq!(tokens, 2);
8288        assert_eq!(window, Some(2));
8289    }
8290
8291    #[test]
8292    fn mla_fixture_emits_latent_state_and_sparse_overflow_refuses() {
8293        use memra_gguf::model_plan::{
8294            MlaAttentionPlan, RopeFactors, RopePlan, SparseIndexPlan, StatePlan,
8295        };
8296
8297        let config = ModelConfig::from_hf(&HfConfig::parse(
8298            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":8,
8299            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8300            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32}"#,
8301        ));
8302        let mut plan = ModelPlan::compile(&config).unwrap();
8303        let mla = MlaAttentionPlan::LatentKv {
8304            query_heads: 2,
8305            q_lora_rank: 4,
8306            kv_lora_rank: 4,
8307            qk_head_dim: 4,
8308            rope_head_dim: 2,
8309            value_head_dim: 4,
8310            rope: RopePlan {
8311                dimensions: 2,
8312                base: 10_000.0,
8313                factors: RopeFactors::None,
8314            },
8315            sparse_index: SparseIndexPlan::None,
8316        };
8317        plan.layers[0].attention = AttentionPlan::Mla(mla.clone());
8318        plan.layers[0].state = StatePlan::LatentKvCache {
8319            width: 6,
8320            index_width: 0,
8321        };
8322        let fixture = deterministic_fixture(&plan).unwrap();
8323        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8324        let ReferenceLayerState::LatentKv { tokens, width, .. } = output.state.layers[0] else {
8325            panic!("expected latent KV state");
8326        };
8327        assert_eq!((tokens, width), (3, 6));
8328        assert_eq!(
8329            output.logits[..4]
8330                .iter()
8331                .map(|value| value.to_bits())
8332                .collect::<Vec<_>>(),
8333            vec![1_035_177_220, 1_055_447_641, 3_201_478_680, 3_199_508_856]
8334        );
8335
8336        let MlaAttentionPlan::LatentKv {
8337            query_heads,
8338            q_lora_rank,
8339            kv_lora_rank,
8340            qk_head_dim,
8341            rope_head_dim,
8342            value_head_dim,
8343            rope,
8344            ..
8345        } = mla
8346        else {
8347            unreachable!()
8348        };
8349        plan.layers[0].attention = AttentionPlan::Mla(MlaAttentionPlan::LatentKv {
8350            query_heads,
8351            q_lora_rank,
8352            kv_lora_rank,
8353            qk_head_dim,
8354            rope_head_dim,
8355            value_head_dim,
8356            rope,
8357            sparse_index: SparseIndexPlan::Own {
8358                heads: 1,
8359                head_dim: 2,
8360                top_k: 2,
8361                kpool: None,
8362            },
8363        });
8364        let error = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap_err();
8365        assert!(matches!(
8366            error,
8367            ReferenceError::UnsupportedOperation {
8368                operation: "sparse MLA selection beyond full-selection equivalence",
8369                ..
8370            }
8371        ));
8372    }
8373
8374    #[test]
8375    fn compressed_mla_executes_window_compressor_indexer_and_grouped_output() {
8376        use memra_gguf::model_plan::{
8377            KvCompressorPlan, MlaAttentionPlan, RopeFactors, RopePlan, SparseIndexPlan, StatePlan,
8378        };
8379
8380        let config = ModelConfig::from_hf(&HfConfig::parse(
8381            r#"{"model_type":"qwen3","num_hidden_layers":1,"hidden_size":128,
8382            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":64,
8383            "intermediate_size":256,"vocab_size":32,"max_position_embeddings":64,
8384            "rms_norm_eps":0.000001}"#,
8385        ));
8386        let mut plan = ModelPlan::compile(&config).unwrap();
8387        plan.layers[0].attention = AttentionPlan::Mla(MlaAttentionPlan::CompressedKv {
8388            query_heads: 2,
8389            q_lora_rank: 64,
8390            latent_head_dim: 128,
8391            rope_head_dim: 64,
8392            output_lora_rank: 64,
8393            output_groups: 1,
8394            window: 4,
8395            rope: RopePlan {
8396                dimensions: 64,
8397                base: 160_000.0,
8398                factors: RopeFactors::Yarn {
8399                    factor: 2.0,
8400                    original_context: 32,
8401                    beta_fast: 32.0,
8402                    beta_slow: 1.0,
8403                },
8404            },
8405            compressor: Some(KvCompressorPlan {
8406                ratio: 4,
8407                latent_dim: 256,
8408            }),
8409            sparse_index: SparseIndexPlan::Own {
8410                heads: 2,
8411                head_dim: 128,
8412                top_k: 2,
8413                kpool: None,
8414            },
8415        });
8416        plan.layers[0].state = StatePlan::CompressedAttention {
8417            window: 4,
8418            head_dim: 128,
8419            compressor_ratio: Some(4),
8420            sparse_top_k: Some(2),
8421        };
8422        let fixture = deterministic_fixture(&plan).unwrap();
8423        let output = execute(&plan, &fixture.weights, &[1, 2, 3, 4]).unwrap();
8424        let ReferenceLayerState::CompressedAttention {
8425            tokens,
8426            width,
8427            window,
8428            compressed_tokens,
8429            ..
8430        } = output.state.layers[0]
8431        else {
8432            panic!("expected compressed attention state")
8433        };
8434        assert_eq!((tokens, width, window, compressed_tokens), (5, 128, 4, 1));
8435        assert!(output.logits.iter().all(|value| value.is_finite()));
8436    }
8437
8438    #[test]
8439    fn dsv4_shaped_trunk_executes_one_canonical_plan() {
8440        let config = ModelConfig::from_hf(&HfConfig::parse(
8441            r#"{"model_type":"deepseek_v4","num_hidden_layers":2,"hidden_size":128,
8442            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":128,
8443            "intermediate_size":256,"vocab_size":128,"max_position_embeddings":1024,
8444            "rms_norm_eps":0.000001,"rope_theta":10000,"n_routed_experts":4,
8445            "n_shared_experts":1,"num_experts_per_tok":2,"moe_intermediate_size":128,
8446            "norm_topk_prob":true,"num_hash_layers":1,"num_nextn_predict_layers":1,
8447            "scoring_func":"sqrtsoftplus","topk_method":"noaux_tc",
8448            "routed_scaling_factor":1.5,"hc_eps":0.000001,"hc_mult":2,
8449            "hc_sinkhorn_iters":4,"q_lora_rank":128,"qk_rope_head_dim":64,
8450            "o_lora_rank":128,"o_groups":1,"index_n_heads":1,"index_head_dim":128,
8451            "index_topk":16,"compress_ratios":[0,4,0],"compress_rope_theta":160000,
8452            "sliding_window":128,"swiglu_limit":10.0,
8453            "rope_scaling":{"factor":4,"beta_fast":32,"beta_slow":1,
8454            "original_max_position_embeddings":1024}}"#,
8455        ));
8456        let mut plan = ModelPlan::compile(&config).unwrap();
8457        assert_eq!(plan.layers.len(), 2);
8458        plan.mtp_blocks.clear();
8459        let fixture = deterministic_fixture(&plan).unwrap();
8460        let output = execute(&plan, &fixture.weights, &[1, 2, 3, 4]).unwrap();
8461        assert_eq!(output.state.layers.len(), 2);
8462        assert!(
8463            output
8464                .state
8465                .layers
8466                .iter()
8467                .all(|state| matches!(state, ReferenceLayerState::CompressedAttention { .. }))
8468        );
8469        assert!(
8470            fixture
8471                .weights
8472                .contains_key(&layer_id(0, LayerTensor::MoeTokenToExpert))
8473        );
8474        assert!(
8475            fixture
8476                .weights
8477                .contains_key(&layer_id(1, LayerTensor::MoeRouterBias))
8478        );
8479        assert!(output.logits.iter().all(|value| value.is_finite()));
8480    }
8481
8482    #[test]
8483    fn dspark_executes_trunk_tap_ring_blocks_markov_and_confidence() {
8484        use memra_gguf::model_plan::{DrafterPlan, DsparkPlan};
8485
8486        let config = ModelConfig::from_hf(&HfConfig::parse(
8487            r#"{"model_type":"deepseek_v4","num_hidden_layers":2,"hidden_size":128,
8488            "num_attention_heads":1,"num_key_value_heads":1,"head_dim":128,
8489            "intermediate_size":256,"vocab_size":128,"max_position_embeddings":1024,
8490            "rms_norm_eps":0.000001,"rope_theta":10000,"n_routed_experts":4,
8491            "n_shared_experts":1,"num_experts_per_tok":2,"moe_intermediate_size":128,
8492            "norm_topk_prob":true,"num_hash_layers":1,"num_nextn_predict_layers":1,
8493            "scoring_func":"sqrtsoftplus","topk_method":"noaux_tc",
8494            "routed_scaling_factor":1.5,"hc_eps":0.000001,"hc_mult":2,
8495            "hc_sinkhorn_iters":4,"q_lora_rank":128,"qk_rope_head_dim":64,
8496            "o_lora_rank":128,"o_groups":1,"index_n_heads":1,"index_head_dim":128,
8497            "index_topk":16,"compress_ratios":[0,4,0],"compress_rope_theta":160000,
8498            "sliding_window":128,"swiglu_limit":10.0,
8499            "rope_scaling":{"factor":4,"beta_fast":32,"beta_slow":1,
8500            "original_max_position_embeddings":1024}}"#,
8501        ));
8502        let mut plan = ModelPlan::compile(&config).unwrap();
8503        let block = plan.mtp_blocks.remove(0).layer;
8504        plan.drafter = Some(DrafterPlan::Dspark(DsparkPlan {
8505            block_size: 3,
8506            noise_token_id: 31,
8507            target_layer_ids: vec![1],
8508            markov_rank: 8,
8509            blocks: vec![block],
8510        }));
8511        let fixture = deterministic_fixture(&plan).unwrap();
8512        let output = execute(&plan, &fixture.weights, &[1, 2, 3, 4]).unwrap();
8513        let draft = output.draft.expect("DSpark output");
8514        assert_eq!(draft.input_token, 4);
8515        assert_eq!(draft.output_ids.len(), 4);
8516        assert_eq!(draft.confidence.len(), 3);
8517        assert_eq!(draft.logits.len(), 3 * 128);
8518        assert!(draft.logits.iter().all(|value| value.is_finite()));
8519        assert!(draft.confidence.iter().all(|value| value.is_finite()));
8520    }
8521
8522    #[test]
8523    fn gemma4_vision_executes_patch_rope_pool_standardize_and_projection() {
8524        let config = ModelConfig::from_hf(&HfConfig::parse(
8525            r#"{"model_type":"gemma4","image_token_id":31,"vision_soft_tokens_per_image":1,
8526            "text_config":{"model_type":"gemma4_text",
8527            "num_hidden_layers":2,"hidden_size":8,"num_attention_heads":2,
8528            "num_key_value_heads":1,"num_global_key_value_heads":1,"head_dim":4,
8529            "global_head_dim":4,"intermediate_size":16,"vocab_size":32,
8530            "max_position_embeddings":64,"rms_norm_eps":0.000001,"sliding_window":8,
8531            "layer_types":["sliding_attention","full_attention"],
8532            "rope_parameters":{"full_attention":{"rope_theta":10000,
8533            "partial_rotary_factor":0.5},"sliding_attention":{"rope_theta":10000}}},
8534            "vision_config":{"hidden_size":8,"intermediate_size":16,
8535            "num_hidden_layers":2,"num_attention_heads":2,"num_key_value_heads":1,
8536            "head_dim":4,"max_position_embeddings":64,"patch_size":2,
8537            "position_embedding_size":16,"pooling_kernel_size":2,
8538            "rms_norm_eps":0.000001,"standardize":true,"use_clipped_linears":false,
8539            "hidden_activation":"gelu_pytorch_tanh","rope_parameters":{"rope_theta":100}}}"#,
8540        ));
8541        let plan = ModelPlan::compile(&config).unwrap();
8542        let fixture = deterministic_fixture(&plan).unwrap();
8543        let input = fixture.vision.as_ref().expect("vision fixture");
8544        let first = execute_vision(&plan, &fixture.weights, input).unwrap();
8545        let second = execute_vision(&plan, &fixture.weights, input).unwrap();
8546        assert_eq!(first, second);
8547        assert_eq!((first.patch_count, first.output_tokens), (4, 1));
8548        assert_eq!((first.hidden_size, first.projection_size), (8, 8));
8549        assert_eq!(first.encoder_hidden.len(), 4 * 8);
8550        assert_eq!(first.pooled_hidden.len(), 8);
8551        assert_eq!(first.projected_hidden.len(), 8);
8552        assert!(first.projected_hidden.iter().all(|value| value.is_finite()));
8553        let multimodal = execute_multimodal(&plan, &fixture.weights, &[1, 31, 2], input).unwrap();
8554        let text_only = execute(&plan, &fixture.weights, &[1, 31, 2]).unwrap();
8555        assert_eq!(multimodal.vision, first);
8556        assert_ne!(multimodal.language.logits, text_only.logits);
8557        assert!(
8558            plan.operations()
8559                .contains(&memra_gguf::model_plan::OperationKind::VisionTokenInjection)
8560        );
8561    }
8562
8563    #[test]
8564    fn gemma4_parallel_moe_executes_shared_routed_and_scaled_residual_branches() {
8565        let config = ModelConfig::from_hf(&HfConfig::parse(
8566            r#"{"model_type":"gemma4","text_config":{"model_type":"gemma4_text",
8567            "num_hidden_layers":2,"hidden_size":8,"num_attention_heads":2,
8568            "num_key_value_heads":1,"num_global_key_value_heads":1,"head_dim":4,
8569            "global_head_dim":4,"intermediate_size":16,"moe_intermediate_size":8,
8570            "num_experts":4,"top_k_experts":2,"vocab_size":32,
8571            "max_position_embeddings":64,"rms_norm_eps":0.000001,"sliding_window":8,
8572            "layer_types":["sliding_attention","full_attention"],
8573            "rope_parameters":{"full_attention":{"rope_theta":10000,
8574            "partial_rotary_factor":0.5},"sliding_attention":{"rope_theta":10000}}}}"#,
8575        ));
8576        let plan = ModelPlan::compile(&config).unwrap();
8577        let MlpPlan::Moe(moe) = &plan.layers[0].mlp else {
8578            panic!("expected Gemma MoE")
8579        };
8580        assert_eq!(moe.experts_per_token, 2);
8581        assert_eq!(moe.shared.as_ref().unwrap().intermediate_size, 16);
8582        assert!(matches!(
8583            plan.layers[0].residual,
8584            ResidualTopology::Gemma {
8585                parallel_moe: Some(_),
8586                ..
8587            }
8588        ));
8589        let fixture = deterministic_fixture(&plan).unwrap();
8590        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8591        assert!(output.logits.iter().all(|value| value.is_finite()));
8592        assert!(
8593            plan.operations()
8594                .contains(&memra_gguf::model_plan::OperationKind::GemmaParallelMoeResidual)
8595        );
8596    }
8597
8598    #[test]
8599    fn embedded_mtp_executes_typed_fusion_block_and_fallback_head() {
8600        let config = ModelConfig::from_hf(&HfConfig::parse(
8601            r#"{"model_type":"qwen3_5","num_hidden_layers":2,
8602            "num_nextn_predict_layers":1,"hidden_size":8,
8603            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8604            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
8605            "rms_norm_eps":0.000001,"full_attention_interval":2,
8606            "linear_conv_kernel_dim":3,"linear_key_head_dim":4,
8607            "linear_value_head_dim":4,"linear_num_key_heads":1,
8608            "linear_num_value_heads":2}"#,
8609        ));
8610        let plan = ModelPlan::compile(&config).unwrap();
8611        assert_eq!(plan.mtp_blocks.len(), 1);
8612        let fixture = deterministic_fixture(&plan).unwrap();
8613        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8614        assert_eq!(output.mtp.len(), 1);
8615        assert_eq!(output.mtp[0].depth, 0);
8616        assert_eq!(output.mtp[0].hidden.len(), fixture.token_ids.len() * 8);
8617        assert_eq!(output.mtp[0].logits.len(), fixture.token_ids.len() * 32);
8618        assert!(output.mtp[0].logits.iter().all(|value| value.is_finite()));
8619        assert_eq!(
8620            output.mtp[0].logits[..4]
8621                .iter()
8622                .map(|value| value.to_bits())
8623                .collect::<Vec<_>>(),
8624            vec![1_042_962_358, 1_044_718_512, 3_171_782_004, 3_189_261_409]
8625        );
8626    }
8627
8628    #[test]
8629    fn multi_depth_mtp_threads_hidden_through_every_typed_block() {
8630        let config = ModelConfig::from_hf(&HfConfig::parse(
8631            r#"{"model_type":"qwen3_5","num_hidden_layers":2,
8632            "num_nextn_predict_layers":2,"hidden_size":8,
8633            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
8634            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
8635            "rms_norm_eps":0.000001,"full_attention_interval":2,
8636            "linear_conv_kernel_dim":3,"linear_key_head_dim":4,
8637            "linear_value_head_dim":4,"linear_num_key_heads":1,
8638            "linear_num_value_heads":2}"#,
8639        ));
8640        let plan = ModelPlan::compile(&config).unwrap();
8641        assert_eq!(plan.mtp_blocks.len(), 2);
8642        let fixture = deterministic_fixture(&plan).unwrap();
8643        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
8644        assert_eq!(
8645            output
8646                .mtp
8647                .iter()
8648                .map(|block| block.depth)
8649                .collect::<Vec<_>>(),
8650            vec![0, 1]
8651        );
8652        assert!(
8653            output
8654                .mtp
8655                .iter()
8656                .flat_map(|block| &block.logits)
8657                .all(|value| value.is_finite())
8658        );
8659        assert_ne!(output.mtp[0].hidden, output.mtp[1].hidden);
8660    }
8661
8662    #[test]
8663    fn rope_uses_neox_split_half_pairs() {
8664        use memra_gguf::model_plan::{RopeFactors, RopePlan};
8665
8666        let mut values = vec![1.0, 2.0, 3.0, 4.0];
8667        apply_rope(&mut values, 1, 1, 4, 4, 10_000.0, None, 1.0);
8668        // Position zero is deliberately unchanged.
8669        assert_eq!(values, vec![1.0, 2.0, 3.0, 4.0]);
8670
8671        let mut values = vec![0.0; 8];
8672        values[4..].copy_from_slice(&[1.0, 2.0, 3.0, 4.0]);
8673        apply_rope(&mut values, 2, 1, 4, 4, 10_000.0, None, 1.0);
8674        let (sin0, cos0) = 1.0f32.sin_cos();
8675        let (sin1, cos1) = 0.01f32.sin_cos();
8676        let row = &values[4..];
8677        assert!((row[0] - (cos0 - 3.0 * sin0)).abs() < 1e-6);
8678        assert!((row[2] - (sin0 + 3.0 * cos0)).abs() < 1e-6);
8679        assert!((row[1] - (2.0 * cos1 - 4.0 * sin1)).abs() < 1e-6);
8680        assert!((row[3] - (2.0 * sin1 + 4.0 * cos1)).abs() < 1e-6);
8681        assert_eq!(
8682            rope_factor_values(
8683                &RopePlan {
8684                    dimensions: 4,
8685                    base: 10_000.0,
8686                    factors: RopeFactors::PartialRotary { factor: 0.5 },
8687                },
8688                &ReferenceWeights::new(),
8689            )
8690            .unwrap(),
8691            (Some(vec![1.0, 1.0e30]), 1.0)
8692        );
8693
8694        // YaRN factors resolve to the transformers-twin divisors + attention factor
8695        // (values pinned in memra-gguf's yarn_divisors test against the banked receipt).
8696        let (yarn_factors, yarn_mscale) = rope_factor_values(
8697            &RopePlan {
8698                dimensions: 4,
8699                base: 10_000.0,
8700                factors: RopeFactors::Yarn {
8701                    factor: 2.0,
8702                    original_context: 8,
8703                    beta_fast: 32.0,
8704                    beta_slow: 1.0,
8705                },
8706            },
8707            &ReferenceWeights::new(),
8708        )
8709        .unwrap();
8710        let yarn_factors = yarn_factors.unwrap();
8711        assert_eq!(yarn_factors[0], 1.0);
8712        assert!((yarn_factors[1] - 2.0).abs() < 1e-6);
8713        assert!((yarn_mscale - 1.069_314_7).abs() < 1e-6);
8714    }
8715
8716    /// Every expected number below is hand-derived from the modular_qwen4_exp.py math
8717    /// (SEMANTICS.md §Gated residual), NOT read back from the code under test.
8718    #[test]
8719    fn gated_residual_read_and_write_match_hand_derived_two_stream_toy() {
8720        let (streams, hidden, rank, tokens) = (2usize, 2usize, 1usize, 1usize);
8721        let wide = streams * hidden;
8722        let prefix = "trunk.layers.0.";
8723        let sublayer = "attn_hyper_connection.";
8724        let insert =
8725            |weights: &mut ReferenceWeights, suffix: &str, shape: &[usize], data: &[f32]| {
8726                weights.insert(
8727                    qwen4exp_family_id(format!("{prefix}{sublayer}{suffix}")),
8728                    weight(shape, data),
8729                );
8730            };
8731        // x = [3,4 | 6,8]: both stream groups are parallel, so grouped normalization maps
8732        // them to the SAME direction — n = (3,4)/sqrt(12.5+1e-6) per group. That equality
8733        // is itself the group-independence assertion.
8734        let x = [3.0, 4.0, 6.0, 8.0];
8735
8736        // Case A: zero down/up/inject weights => w = sigmoid(0) = 0.5 everywhere and
8737        // inject = 2*sigmoid(0) = 1; mixed[c] = 0.5*(n0[c]+n1[c])/2.
8738        let mut weights = ReferenceWeights::new();
8739        insert(&mut weights, "hc_norm.weight", &[wide], &[1.0; 4]);
8740        insert(
8741            &mut weights,
8742            "input_mix_weight_down.weight",
8743            &[rank, wide],
8744            &[0.0; 4],
8745        );
8746        insert(
8747            &mut weights,
8748            "input_mix_weight_up.weight",
8749            &[wide, rank],
8750            &[0.0; 4],
8751        );
8752        insert(
8753            &mut weights,
8754            "block_inject_weight.weight",
8755            &[streams, wide],
8756            &[0.0; 8],
8757        );
8758        let (mixed, inject) = gated_residual_read(
8759            &weights, prefix, sublayer, &x, tokens, streams, hidden, rank, 1e-6, true,
8760        )
8761        .unwrap();
8762        // Hand: n = (0.84852810, 1.13137080); mixed = (0.42426406, 0.56568541).
8763        assert!((mixed[0] - 0.424_264_06).abs() < 1e-5, "{mixed:?}");
8764        assert!((mixed[1] - 0.565_685_41).abs() < 1e-5, "{mixed:?}");
8765        assert!((inject[0] - 1.0).abs() < 1e-6 && (inject[1] - 1.0).abs() < 1e-6);
8766
8767        // Case B: down = [1,0,0,0], up = ones, inject row0 = [1,0,0,0], row1 = 0.
8768        // low  = silu(n[0]/2)           = silu(0.42426406)   = 0.25646896
8769        // w    = sigmoid(low)           = 0.56376809 for every dim
8770        // mixed[c] = w*(n0[c]+n1[c])/2  = (0.47837307, 0.63783076)
8771        // inject   = (2*sigmoid(n[0]/2), 2*sigmoid(0)) = (1.20900630, 1.0)
8772        insert(
8773            &mut weights,
8774            "input_mix_weight_down.weight",
8775            &[rank, wide],
8776            &[1.0, 0.0, 0.0, 0.0],
8777        );
8778        insert(
8779            &mut weights,
8780            "input_mix_weight_up.weight",
8781            &[wide, rank],
8782            &[1.0; 4],
8783        );
8784        insert(
8785            &mut weights,
8786            "block_inject_weight.weight",
8787            &[streams, wide],
8788            &[1.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0],
8789        );
8790        let (mixed, inject) = gated_residual_read(
8791            &weights, prefix, sublayer, &x, tokens, streams, hidden, rank, 1e-6, true,
8792        )
8793        .unwrap();
8794        assert!((mixed[0] - 0.478_373_07).abs() < 1e-5, "{mixed:?}");
8795        assert!((mixed[1] - 0.637_830_76).abs() < 1e-5, "{mixed:?}");
8796        assert!((inject[0] - 1.208_999_4).abs() < 1e-4, "{inject:?}");
8797        assert!((inject[1] - 1.0).abs() < 1e-6, "{inject:?}");
8798
8799        // Write: out = PRE-norm wide + block_out ⊗ inject, block_out = (1, -1)
8800        // => (3+1.209, 4-1.209, 6+1, 8-1).
8801        let mut wide_state = x.to_vec();
8802        gated_residual_write(
8803            &mut wide_state,
8804            &[1.0, -1.0],
8805            &inject,
8806            tokens,
8807            streams,
8808            hidden,
8809        );
8810        assert!((wide_state[0] - 4.209_006_3).abs() < 1e-4, "{wide_state:?}");
8811        assert!((wide_state[1] - 2.790_993_7).abs() < 1e-4, "{wide_state:?}");
8812        assert!((wide_state[2] - 7.0).abs() < 1e-6, "{wide_state:?}");
8813        assert!((wide_state[3] - 7.0).abs() < 1e-6, "{wide_state:?}");
8814    }
8815
8816    /// GDN with the qwen4_exp sigmoid z-gate — the ONE divergence from qwen3_5. Single
8817    /// token, identity-shaped projections, gate logit 2.0:
8818    ///   conv (k=1, w=1) => q=k=(silu(1),0), v=(silu(2),0); l2norm makes q~=k unit;
8819    ///   beta=sigmoid(0)=0.5, one step from zero state => mixed = (k.q)*v*beta/sqrt(2);
8820    ///   rms_norm => (1.41420992, 0); out = norm * act(2).
8821    /// Hand: sigmoid arm (1.24563196, 0); silu arm would be (2.49126392, 0).
8822    #[test]
8823    fn gdn_sigmoid_gate_matches_hand_derived_single_token() {
8824        use memra_gguf::model_plan::GatedDeltaNetPlan;
8825
8826        let hidden = 2usize;
8827        let mut weights = ReferenceWeights::new();
8828        weights.insert(
8829            layer_id(0, LayerTensor::GdnQkv),
8830            weight(
8831                &[6, 2],
8832                &[1.0, 0.0, 0.0, 1.0, 1.0, 0.0, 0.0, 1.0, 2.0, 0.0, 0.0, 2.0],
8833            ),
8834        );
8835        weights.insert(
8836            layer_id(0, LayerTensor::GdnGate),
8837            weight(&[2, 2], &[2.0, 0.0, 0.0, 1.0]),
8838        );
8839        weights.insert(
8840            layer_id(0, LayerTensor::GdnBeta),
8841            weight(&[1, 2], &[0.0, 0.0]),
8842        );
8843        weights.insert(
8844            layer_id(0, LayerTensor::GdnAlpha),
8845            weight(&[1, 2], &[0.0, 0.0]),
8846        );
8847        weights.insert(layer_id(0, LayerTensor::GdnA), weight(&[1], &[0.0]));
8848        weights.insert(layer_id(0, LayerTensor::GdnDtBias), weight(&[1], &[0.0]));
8849        weights.insert(layer_id(0, LayerTensor::GdnNorm), weight(&[2], &[1.0, 1.0]));
8850        weights.insert(
8851            layer_id(0, LayerTensor::GdnConv1d),
8852            weight(&[6, 1], &[1.0; 6]),
8853        );
8854        weights.insert(
8855            layer_id(0, LayerTensor::GdnOutput),
8856            weight(&[2, 2], &[1.0, 0.0, 0.0, 1.0]),
8857        );
8858        let plan = GatedDeltaNetPlan {
8859            key_heads: 1,
8860            value_heads: 1,
8861            key_head_dim: 2,
8862            value_head_dim: 2,
8863            conv_kernel: 1,
8864            gate_activation: GdnGateActivation::Sigmoid,
8865        };
8866        let (sigmoid_out, _) =
8867            gated_delta_net(0, &plan, 1e-6, &weights, &[1.0, 0.0], 1, hidden).unwrap();
8868        assert!(
8869            (sigmoid_out[0] - 1.245_632_0).abs() < 1e-4,
8870            "{sigmoid_out:?}"
8871        );
8872        assert!(sigmoid_out[1].abs() < 1e-6, "{sigmoid_out:?}");
8873
8874        let silu_plan = GatedDeltaNetPlan {
8875            gate_activation: GdnGateActivation::Silu,
8876            ..plan
8877        };
8878        let (silu_out, _) =
8879            gated_delta_net(0, &silu_plan, 1e-6, &weights, &[1.0, 0.0], 1, hidden).unwrap();
8880        assert!((silu_out[0] - 2.491_263_9).abs() < 1e-4, "{silu_out:?}");
8881    }
8882
8883    /// Crafted 12-token sequence with an unambiguous top-k choice: only tokens 4..8 carry
8884    /// key mass (block 1); every other block pools to the ZERO vector, which rope and
8885    /// normalization preserve, so its relu score is exactly 0 while block 1 scores
8886    /// strictly positive (2*cos(Δpos) with Δpos ∈ {5, 7} rad, both cos > 0). budget = 1
8887    /// block. Also pins the tail rule, including the boundary case where a query's own
8888    /// unselected complete block makes the query NOT see itself.
8889    #[test]
8890    fn micro_block_indexer_selects_unambiguous_block_and_always_keeps_the_tail() {
8891        let tokens = 12usize;
8892        let hidden = 2usize;
8893        let overlay = MicroBlockIndexPlan {
8894            query_heads: 1,
8895            kv_heads: 1,
8896            head_dim: 2,
8897            rope_dimensions: 2,
8898            block_size: 4,
8899            budget_blocks: 1,
8900            budget_tokens: 4,
8901        };
8902        let rope = RopePlan {
8903            dimensions: 2,
8904            base: 10_000.0,
8905            factors: memra_gguf::model_plan::RopeFactors::None,
8906        };
8907        let prefix = "trunk.layers.0.";
8908        let mut weights = ReferenceWeights::new();
8909        // q rows = identity (q = x); k rows read only x[1] scaled by 10.
8910        weights.insert(
8911            qwen4exp_family_id(format!("{prefix}self_attn.indexer.index_qk_proj.weight")),
8912            weight(&[4, 2], &[1.0, 0.0, 0.0, 1.0, 0.0, 10.0, 0.0, 0.0]),
8913        );
8914        for norm in ["q_layernorm", "k_layernorm"] {
8915            weights.insert(
8916                qwen4exp_family_id(format!("{prefix}self_attn.indexer.{norm}.weight")),
8917                weight(&[2], &[1.0, 1.0]),
8918            );
8919        }
8920        let mut x = vec![0.0; tokens * hidden];
8921        for token in 0..tokens {
8922            x[token * hidden] = 1.0; // every query is (1, 0)
8923            if (4..8).contains(&token) {
8924                x[token * hidden + 1] = 1.0; // block-1 keys become (10, 0)
8925            }
8926        }
8927        let mask = micro_block_selection_mask(
8928            0, &overlay, &rope, 1e-6, &weights, prefix, &x, tokens, hidden,
8929        )
8930        .unwrap();
8931        let row = |token: usize| &mask[token * tokens..(token + 1) * tokens];
8932        // t=0: no complete block, tail = {0}.
8933        assert_eq!(
8934            row(0),
8935            &[
8936                true, false, false, false, false, false, false, false, false, false, false, false
8937            ]
8938        );
8939        // t=5: one complete block (0..4, the only candidate) + tail {4,5}.
8940        assert_eq!(
8941            row(5),
8942            &[
8943                true, true, true, true, true, true, false, false, false, false, false, false
8944            ]
8945        );
8946        // t=9: blocks {0,1} complete, block 1 wins (score>0 vs 0), tail {8,9}.
8947        assert_eq!(
8948            row(9),
8949            &[
8950                false, false, false, false, true, true, true, true, true, true, false, false
8951            ]
8952        );
8953        // t=11: blocks {0,1,2} complete, tail EMPTY; only block 1 selected — the query
8954        // does not even see itself (blocks-only selection at exact boundaries).
8955        assert_eq!(
8956            row(11),
8957            &[
8958                false, false, false, false, true, true, true, true, false, false, false, false
8959            ]
8960        );
8961    }
8962
8963    /// The selection mask gates full attention to exactly the selected sources: a query
8964    /// restricted to itself must return its own VALUE row bit-for-bit reasoning
8965    /// (softmax over one score = 1), with identity projections output == input.
8966    #[test]
8967    fn full_attention_selection_mask_restricts_sources_to_hand_derived_rows() {
8968        use memra_gguf::model_plan::{FullAttentionPlan, RopeFactors, TensorPresence};
8969
8970        let plan = FullAttentionPlan {
8971            query_heads: 1,
8972            kv_heads: 1,
8973            key_head_dim: 2,
8974            value_head_dim: 2,
8975            rope: RopePlan {
8976                dimensions: 2,
8977                base: 10_000.0,
8978                factors: RopeFactors::None,
8979            },
8980            qk_norm: TensorPresence::Absent,
8981            output_gate: memra_gguf::config::AttentionGateKind::None,
8982            scale: AttentionScale::InverseSqrtKeyDim,
8983            value_projection: ValueProjection::Separate,
8984            value_norm: ValueNorm::None,
8985        };
8986        let identity = [1.0, 0.0, 0.0, 1.0];
8987        let mut weights = ReferenceWeights::new();
8988        for tensor in [
8989            LayerTensor::Query,
8990            LayerTensor::Key,
8991            LayerTensor::Value,
8992            LayerTensor::AttentionOutput,
8993        ] {
8994            weights.insert(layer_id(0, tensor), weight(&[2, 2], &identity));
8995        }
8996        // Orthogonal rows keep the causal softmax far from saturation (score gap ~1.3),
8997        // so the unmasked row visibly mixes ~21% of source 0.
8998        let x = [1.0, 0.0, 0.0, 1.0];
8999        let diagonal = [true, false, false, true];
9000        let (masked, _) =
9001            full_attention(0, &plan, None, 1e-6, &weights, &x, 2, 2, Some(&diagonal)).unwrap();
9002        for index in 0..4 {
9003            assert!((masked[index] - x[index]).abs() < 1e-6, "{masked:?}");
9004        }
9005        let (unmasked, _) = full_attention(0, &plan, None, 1e-6, &weights, &x, 2, 2, None).unwrap();
9006        assert!(
9007            (unmasked[2] - x[2]).abs() > 1e-3,
9008            "causal row must mix sources"
9009        );
9010
9011        let starving = [true, false, false, false];
9012        let error =
9013            full_attention(0, &plan, None, 1e-6, &weights, &x, 2, 2, Some(&starving)).unwrap_err();
9014        assert!(matches!(error, ReferenceError::InvalidPlan { .. }));
9015    }
9016
9017    /// N-gram id math recomputed independently below (wrapping i64 multiply, XOR, floor
9018    /// mod, offset — SEMANTICS.md §PLE), with a multiplier big enough that the product
9019    /// wraps negative and exercises the floor-mod arm.
9020    #[test]
9021    fn ngram_ids_match_independently_computed_hash_chain() {
9022        let multipliers = [0x4000_0000_0000_0001_i64, 1_000_003, 7_777_777];
9023        let sizes = [97_i64, 89, 83, 79];
9024        let offsets = [0_i64, 97, 186, 269];
9025        let (max_ngram, heads_per_ngram, eos) = (3usize, 2usize, 9u32);
9026        let token_ids = [5u32, 7];
9027        let ids = ngram_ids(
9028            &token_ids,
9029            &multipliers,
9030            &sizes,
9031            &offsets,
9032            max_ngram,
9033            heads_per_ngram,
9034            eos,
9035            0,
9036        )
9037        .unwrap();
9038
9039        // history = [9, 9, 5, 7]; shifted[1] = [9,9,9,5]; shifted[2] = [9,9,9,9]
9040        // (the two context positions read EOS; position 3 shifted-by-1 reads token 5).
9041        let expect = |mixed: i64, head: usize| mixed.rem_euclid(sizes[head]) + offsets[head];
9042        let bigram_t0 = 5_i64.wrapping_mul(multipliers[0]) ^ 9_i64.wrapping_mul(multipliers[1]);
9043        let trigram_t0 = bigram_t0 ^ 9_i64.wrapping_mul(multipliers[2]);
9044        let bigram_t1 = 7_i64.wrapping_mul(multipliers[0]) ^ 5_i64.wrapping_mul(multipliers[1]);
9045        let trigram_t1 = bigram_t1 ^ 9_i64.wrapping_mul(multipliers[2]);
9046        // 7 * (2^62 + 1) wraps to 2^63 + 2^62 + 7, i.e. negative i64; floor mod must
9047        // still land non-negative (torch.remainder semantics).
9048        assert!(7_i64.wrapping_mul(multipliers[0]) < 0);
9049        assert_eq!(
9050            ids,
9051            vec![
9052                expect(bigram_t0, 0),
9053                expect(bigram_t0, 1),
9054                expect(trigram_t0, 2),
9055                expect(trigram_t0, 3),
9056                expect(bigram_t1, 0),
9057                expect(bigram_t1, 1),
9058                expect(trigram_t1, 2),
9059                expect(trigram_t1, 3),
9060            ]
9061        );
9062        assert!(ids.iter().all(|&id| id >= 0));
9063    }
9064
9065    /// Hand-derived shift vectors for history [E,E,5,6,E,7,8] (E = 63):
9066    ///   eos strictly-before: [-1,0,1,1,1,4,4]; segment starts [0,1,2,2,2,5,5];
9067    ///   in-segment positions [0,0,0,1,2,0,1].
9068    /// shift=1 keeps positions {3,4,6} (note position 4 — the EOS itself — reads 6, its
9069    /// in-segment index counts within the PREVIOUS segment); shift=2 keeps only {4}.
9070    #[test]
9071    fn eos_segment_reset_reads_eos_across_boundaries() {
9072        let eos = 63i64;
9073        let history = [eos, eos, 5, 6, eos, 7, 8];
9074        assert_eq!(shift_right_ignore_eos(&history, 0, eos), history.to_vec());
9075        assert_eq!(
9076            shift_right_ignore_eos(&history, 1, eos),
9077            vec![eos, eos, eos, 5, 6, eos, 7]
9078        );
9079        assert_eq!(
9080            shift_right_ignore_eos(&history, 2, eos),
9081            vec![eos, eos, eos, eos, 5, eos, eos]
9082        );
9083    }
9084
9085    /// Scalar-channel PLE block pinning the gather -> gate -> dilated-conv chain by hand:
9086    /// wide stream 0 => query norm 0 => gate = sigmoid(0) = 0.5, so gated = 0.5*value;
9087    /// normed scalars n_t = g_t/sqrt(g_t^2+1e-6); conv (kernel 2, dilation = max_ngram
9088    /// = 2, taps w = [10, 1]) reads out[t] = g_t + silu(10*n_{t-2} + n_t) with the
9089    /// out-of-range tap dropped. Hand values below; a REVERSED tap order would give
9090    /// out[2] = 11.5068593 instead of 8.8685460, so this pins conv orientation AND the
9091    /// dilation reach (t-2, not t-1).
9092    #[test]
9093    fn ple_block_matches_hand_derived_scalar_gather_gate_and_dilated_conv() {
9094        let prefix = "trunk.layers.1.";
9095        let mut weights = ReferenceWeights::new();
9096        let family = |suffix: &str| qwen4exp_family_id(format!("{prefix}{suffix}"));
9097        weights.insert(
9098            family("ple.ple_embedding.layer_multipliers"),
9099            ReferenceTensor::new_i64(vec![2], vec![1, 0]).unwrap(),
9100        );
9101        weights.insert(
9102            family("ple.ple_embedding.ngram_heads_vocab_sizes"),
9103            ReferenceTensor::new_i64(vec![1], vec![5]).unwrap(),
9104        );
9105        weights.insert(
9106            family("ple.ple_embedding.ngram_heads_offsets"),
9107            ReferenceTensor::new_i64(vec![1], vec![0]).unwrap(),
9108        );
9109        // ids = token mod 5 = [1, 2, 3] -> values [0.002, 0.4, 1.6]
9110        weights.insert(
9111            family("ple.ple_embedding.ngram_embedding"),
9112            weight(&[5, 1], &[0.0, 0.002, 0.4, 1.6, 0.0]),
9113        );
9114        weights.insert(family("ple.key_proj.weight"), weight(&[1, 1], &[1.0]));
9115        weights.insert(family("ple.value_proj.weight"), weight(&[1, 1], &[1.0]));
9116        for norm in ["norm_key", "norm_query", "norm_conv"] {
9117            weights.insert(family(&format!("ple.{norm}.weight")), weight(&[1], &[1.0]));
9118        }
9119        weights.insert(family("ple.conv1d.weight"), weight(&[1, 2], &[10.0, 1.0]));
9120        let plan = memra_gguf::model_plan::PleEmbeddingPlan {
9121            ngram_heads: 1,
9122            head_embed_dim: 1,
9123            vocab_shards: 1,
9124            embed_dim: 1,
9125            conv_kernel: 2,
9126            max_ngram: 2,
9127            eos_token_id: 4,
9128        };
9129        let wide_state = [0.0; 3];
9130        let output = ple_block(
9131            1,
9132            &plan,
9133            1e-6,
9134            &weights,
9135            prefix,
9136            &wide_state,
9137            &[1, 2, 3],
9138            3,
9139            1,
9140            1,
9141        )
9142        .unwrap();
9143        assert!((output[0] - 0.474_592_9).abs() < 1e-4, "{output:?}");
9144        assert!((output[1] - 0.931_047_0).abs() < 1e-4, "{output:?}");
9145        assert!((output[2] - 8.868_546_0).abs() < 1e-3, "{output:?}");
9146    }
9147
9148    /// The qwen4_exp pack's tiny plan executes end-to-end through `execute`: gated
9149    /// residual entry/exit, GDN + QSA trunk, PLE on layer 1, MoE with the gated shared
9150    /// expert, and the separate-projection MTP block. 16 tokens so the indexer budget
9151    /// (2 blocks) actually BINDS (3-4 complete blocks at the last queries).
9152    #[test]
9153    fn qwen4exp_tiny_plan_executes_gated_residual_qsa_ple_moe_and_mtp() {
9154        let pack = memra_gguf::model_packs::by_alias("qwen4_exp").expect("qwen4_exp pack");
9155        let plan = pack.compile_tiny_plan().expect("tiny plan compiles");
9156        assert_eq!(plan.layers.len(), 4);
9157        assert_eq!(plan.mtp_blocks.len(), 1);
9158        let fixture = deterministic_fixture(&plan).unwrap();
9159        assert!(
9160            !fixture.weights.contains_key(&TensorId::OutputNorm),
9161            "exit-mixer plans must not fabricate a final norm"
9162        );
9163        let token_ids: Vec<u32> = (1..=16).collect();
9164        let first = execute(&plan, &fixture.weights, &token_ids).unwrap();
9165        let second = execute(&plan, &fixture.weights, &token_ids).unwrap();
9166        assert_eq!(first, second, "reference must be bit-deterministic");
9167        assert_eq!((first.tokens, first.vocab), (16, 64));
9168        assert!(first.logits.iter().all(|value| value.is_finite()));
9169        for (index, state) in first.state.layers.iter().enumerate() {
9170            if index == 3 {
9171                assert!(matches!(state, ReferenceLayerState::Kv { .. }));
9172            } else {
9173                assert!(matches!(state, ReferenceLayerState::Recurrent { .. }));
9174            }
9175        }
9176        // MTP: wide (streams*hidden) post-layer state is the K>1 carrier.
9177        assert_eq!(first.mtp.len(), 1);
9178        assert_eq!(first.mtp[0].hidden.len(), 16 * 2 * 16);
9179        assert_eq!(first.mtp[0].logits.len(), 16 * 64);
9180        assert!(first.mtp[0].logits.iter().all(|value| value.is_finite()));
9181
9182        // The QSA selection binds: zeroing the indexer projection makes every block
9183        // score exactly 0, so the pinned tie rule keeps the LOWEST-indexed blocks —
9184        // a different selection than the trained-shaped fixture picks. (A sign flip
9185        // would NOT work here: negating q and k together preserves every score.)
9186        let mut perturbed = fixture.weights.clone();
9187        perturbed
9188            .get_mut(&qwen4exp_family_id(
9189                "trunk.layers.3.self_attn.indexer.index_qk_proj.weight".into(),
9190            ))
9191            .expect("trunk indexer weights")
9192            .data
9193            .fill(0.0);
9194        let reindexed = execute(&plan, &perturbed, &token_ids).unwrap();
9195        assert_ne!(
9196            first.logits, reindexed.logits,
9197            "indexer selection must gate attention"
9198        );
9199
9200        // PLE binds: a different n-gram table moves the logits.
9201        let mut retabled = fixture.weights.clone();
9202        retabled
9203            .get_mut(&qwen4exp_family_id(
9204                "trunk.layers.1.ple.ple_embedding.ngram_embedding".into(),
9205            ))
9206            .expect("ngram table")
9207            .data
9208            .fill(0.25);
9209        let regathered = execute(&plan, &retabled, &token_ids).unwrap();
9210        assert_ne!(
9211            first.logits, regathered.logits,
9212            "PLE gather must feed layer 1"
9213        );
9214
9215        // The sigmoid-gated shared expert binds (MoE deliverable check).
9216        let mut regated = fixture.weights.clone();
9217        regated
9218            .get_mut(&layer_id(0, LayerTensor::SharedMlpInputGate))
9219            .expect("shared expert gate")
9220            .data
9221            .fill(4.0);
9222        let reshared = execute(&plan, &regated, &token_ids).unwrap();
9223        assert_ne!(
9224            first.logits, reshared.logits,
9225            "shared-expert sigmoid gate must scale the shared branch"
9226        );
9227    }
9228
9229    #[test]
9230    fn dense_gemma_executes_scaled_parallel_residual_and_k_as_v() {
9231        let config = ModelConfig::from_hf(&HfConfig::parse(
9232            r#"{"model_type":"gemma4","num_hidden_layers":2,"hidden_size":8,
9233            "num_attention_heads":2,"num_key_value_heads":1,
9234            "num_global_key_value_heads":1,"head_dim":4,"global_head_dim":4,
9235            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
9236            "rms_norm_eps":0.000001,"sliding_window":2,
9237            "final_logit_softcapping":30,
9238            "layer_types":["sliding_attention","full_attention"],
9239            "rope_parameters":{"full_attention":{"rope_theta":1000000,
9240            "partial_rotary_factor":0.5},"sliding_attention":{"rope_theta":10000}}}"#,
9241        ));
9242        let plan = ModelPlan::compile(&config).unwrap();
9243        assert_eq!(plan.embedding_scale, 8.0f32.sqrt());
9244        let fixture = deterministic_fixture(&plan).unwrap();
9245        assert!(
9246            !fixture
9247                .weights
9248                .contains_key(&layer_id(1, LayerTensor::Value))
9249        );
9250        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
9251        assert!(output.logits.iter().all(|value| value.is_finite()));
9252        let ReferenceLayerState::Kv { window, .. } = output.state.layers[0] else {
9253            panic!("expected SWA state");
9254        };
9255        assert_eq!(window, Some(2));
9256        let ReferenceLayerState::Kv { window, .. } = output.state.layers[1] else {
9257            panic!("expected global state");
9258        };
9259        assert_eq!(window, None);
9260        assert_eq!(
9261            output.logits[..4]
9262                .iter()
9263                .map(|value| value.to_bits())
9264                .collect::<Vec<_>>(),
9265            vec![3_198_203_366, 1_057_194_687, 3_185_247_713, 3_204_119_266]
9266        );
9267    }
9268
9269    /// ONE TensorId, ONE byte order. `deterministic_fixture` mints the reference's weights and
9270    /// `TensorContract` names the same ids for the engine's loader; any engine-vs-reference gate
9271    /// serves ONE set of bytes under both. A shape disagreement preserves element counts, so
9272    /// nothing else catches it — `MlaKeyUp` was minted `[head][nope][rank]` against the
9273    /// contract's `[head][rank][nope]` and every MLA parity comparison silently mis-strided the
9274    /// absorb operand (glm53-flash lane, 2026-08-28: 1.24e-1 relative on a micro fixture that
9275    /// drops to 6.9e-7 once the layouts agree).
9276    #[test]
9277    fn the_mla_fixture_shapes_match_the_tensor_contract() {
9278        use memra_gguf::tensor_contract::{
9279            CheckpointDialect, ContractOptions, OutputHead, TensorContract,
9280        };
9281
9282        use memra_gguf::model_plan::{MlaAttentionPlan, StatePlan};
9283
9284        // EVERY MLA extent DISTINCT (heads 2, q_lora 3, kv_rank 6, nope 4, v 5). The shared tiny
9285        // plan runs 2/4/4/4/4, where a transposed key plane has the SAME shape as a correct one
9286        // and this pin would be a tautology.
9287        let mut plan = kpool_mla_reference_plan();
9288        let AttentionPlan::Mla(MlaAttentionPlan::LatentKv {
9289            q_lora_rank,
9290            kv_lora_rank,
9291            qk_head_dim,
9292            value_head_dim,
9293            ..
9294        }) = &mut plan.layers[1].attention
9295        else {
9296            panic!("layer 1 of the tiny plan must be MLA LatentKv");
9297        };
9298        *q_lora_rank = 3;
9299        *kv_lora_rank = 6;
9300        *qk_head_dim = 4;
9301        *value_head_dim = 5;
9302        plan.layers[1].state = StatePlan::LatentKvCache {
9303            width: 6,
9304            index_width: 8,
9305        };
9306        let fixture = deterministic_fixture(&plan).unwrap();
9307        let contract = TensorContract::for_plan(
9308            &plan,
9309            CheckpointDialect::Gguf,
9310            ContractOptions {
9311                output_head: OutputHead::TiedToEmbedding,
9312            },
9313        )
9314        .unwrap();
9315        let mut checked = 0;
9316        for requirement in &contract.requirements {
9317            let Some(tensor) = fixture.weights.get(&requirement.id) else {
9318                continue;
9319            };
9320            // GGUF `ne` is fastest-axis-first; the fixture states row-major shapes.
9321            let mut wanted: Vec<usize> = requirement.shape.iter().map(|&d| d as usize).collect();
9322            wanted.reverse();
9323            let TensorId::Layer { tensor: kind, .. } = requirement.id else {
9324                continue;
9325            };
9326            if !matches!(
9327                kind,
9328                LayerTensor::MlaKeyUp | LayerTensor::MlaValueUp | LayerTensor::MlaQueryUp
9329            ) {
9330                continue;
9331            }
9332            assert_eq!(
9333                tensor.shape, wanted,
9334                "{:?}: fixture shape {:?} but the contract declares ne {:?}",
9335                requirement.id, tensor.shape, requirement.shape
9336            );
9337            checked += 1;
9338        }
9339        assert!(checked >= 3, "the plan must exercise the MLA planes");
9340    }
9341
9342    /// The glm5_next-shaped tiny plan: one KDA layer, one MLA+k-pool-indexer layer, sigmoid MoE,
9343    /// hyper-connections with the mean collapse. Shared by the execution gate and the
9344    /// fixture-vs-contract shape pin so both describe the SAME plan.
9345    fn kpool_mla_reference_plan() -> ModelPlan {
9346        use memra_gguf::model_plan::{
9347            DenseMlpPlan, KimiDeltaNetPlan, KpoolPlan, MlaAttentionPlan, MoeMlpPlan, RopeFactors,
9348            RopePlan, RouterPlan, SharedMlpPlan, SparseIndexPlan, StatePlan,
9349        };
9350
9351        let config = ModelConfig::from_hf(&HfConfig::parse(
9352            r#"{"model_type":"qwen3","num_hidden_layers":2,"hidden_size":8,
9353            "num_attention_heads":2,"num_key_value_heads":1,"head_dim":4,
9354            "intermediate_size":16,"vocab_size":32,"max_position_embeddings":32,
9355            "rms_norm_eps":0.00001}"#,
9356        ));
9357        let mut plan = ModelPlan::compile(&config).unwrap();
9358        plan.layers[0].attention = AttentionPlan::KimiDeltaNet(KimiDeltaNetPlan {
9359            num_heads: 2,
9360            head_dim: 4,
9361            conv_kernel: 3,
9362            gate_lower_bound: -5.0,
9363        });
9364        plan.layers[0].state = StatePlan::Recurrent {
9365            conv_width: 24,
9366            conv_kernel: 3,
9367            state_width: 32,
9368        };
9369        plan.layers[0].mlp = MlpPlan::Dense(DenseMlpPlan {
9370            intermediate_size: 16,
9371            activation: ActivationPlan::SwiGluPreClamped { limit: 10.0 },
9372        });
9373        plan.layers[1].attention = AttentionPlan::Mla(MlaAttentionPlan::LatentKv {
9374            query_heads: 2,
9375            q_lora_rank: 4,
9376            kv_lora_rank: 4,
9377            qk_head_dim: 4,
9378            rope_head_dim: 0,
9379            value_head_dim: 4,
9380            rope: RopePlan {
9381                dimensions: 0,
9382                base: 10_000.0,
9383                factors: RopeFactors::None,
9384            },
9385            sparse_index: SparseIndexPlan::Own {
9386                heads: 2,
9387                head_dim: 4,
9388                top_k: 4,
9389                kpool: Some(KpoolPlan {
9390                    pool: 2,
9391                    always_select_tail: true,
9392                }),
9393            },
9394        });
9395        plan.layers[1].state = StatePlan::LatentKvCache {
9396            width: 4,
9397            index_width: 8,
9398        };
9399        plan.layers[1].mlp = MlpPlan::Moe(MoeMlpPlan {
9400            expert_count: 4,
9401            experts_per_token: 2,
9402            expert_intermediate_size: 4,
9403            router: RouterPlan::Sigmoid {
9404                normalize_selected: true,
9405                scaling_factor: 2.5,
9406                selection_bias: true,
9407            },
9408            shared: Some(SharedMlpPlan {
9409                intermediate_size: 4,
9410                gated: false,
9411            }),
9412            activation: ActivationPlan::SwiGluPreClamped { limit: 10.0 },
9413        });
9414        for layer in &mut plan.layers {
9415            layer.residual = ResidualTopology::HyperConnections {
9416                streams: 2,
9417                epsilon: 1e-6,
9418                sinkhorn_iterations: 2,
9419                collapse: HcCollapse::Mean,
9420            };
9421        }
9422        plan
9423    }
9424
9425    #[test]
9426    fn glm5_shaped_tiny_plan_executes_kda_kpool_mla_and_mean_collapse_deterministically() {
9427        let plan = kpool_mla_reference_plan();
9428        let fixture = deterministic_fixture(&plan).unwrap();
9429        // The mean collapse owns no learned head tensors.
9430        assert!(!fixture.weights.contains_key(&TensorId::HyperHeadFunction));
9431        assert!(
9432            fixture
9433                .weights
9434                .contains_key(&layer_id(1, LayerTensor::SparseCompressorGate))
9435        );
9436        let output = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
9437        assert_eq!(output.logits.len(), fixture.token_ids.len() * 32);
9438        assert!(output.logits.iter().all(|value| value.is_finite()));
9439        assert!(matches!(
9440            output.state.layers[0],
9441            ReferenceLayerState::Recurrent { conv_width: 24, .. }
9442        ));
9443        assert!(matches!(
9444            output.state.layers[1],
9445            ReferenceLayerState::LatentKv { width: 4, .. }
9446        ));
9447        let second = execute(&plan, &fixture.weights, &fixture.token_ids).unwrap();
9448        assert_eq!(
9449            output
9450                .logits
9451                .iter()
9452                .map(|value| value.to_bits())
9453                .collect::<Vec<_>>(),
9454            second
9455                .logits
9456                .iter()
9457                .map(|value| value.to_bits())
9458                .collect::<Vec<_>>()
9459        );
9460    }
9461
9462    #[test]
9463    fn kimi_delta_net_matches_hand_derived_three_token_recurrence() {
9464        use memra_gguf::model_plan::KimiDeltaNetPlan;
9465
9466        let plan = KimiDeltaNetPlan {
9467            num_heads: 1,
9468            head_dim: 2,
9469            conv_kernel: 2,
9470            gate_lower_bound: -5.0,
9471        };
9472        let x = [[0.5f32, -0.3], [0.1, 0.8], [-0.6, 0.2]];
9473        let wq = [[0.7f32, -0.2], [0.3, 0.5]];
9474        let wk = [[0.4f32, 0.1], [-0.3, 0.6]];
9475        let wv = [[0.9f32, 0.2], [-0.1, 0.8]];
9476        let q_conv = [[0.3f32, 0.7], [-0.2, 0.9]];
9477        let k_conv = [[0.5f32, 0.5], [0.1, 0.8]];
9478        let v_conv = [[0.2f32, 0.6], [0.4, 0.4]];
9479        let f_a = [[0.6f32, -0.4], [0.2, 0.3]];
9480        let f_b = [[0.5f32, 0.1], [-0.2, 0.7]];
9481        let dt_bias = [0.05f32, -0.1];
9482        let a_log = [0.2f32];
9483        let b_proj = [[0.4f32, -0.6]];
9484        let g_a = [[0.3f32, 0.2], [-0.5, 0.4]];
9485        let g_b = [[0.6f32, -0.3], [0.2, 0.5]];
9486        let o_norm = [1.0f32, 1.5];
9487        let wo = [[0.8f32, -0.4], [0.3, 0.9]];
9488
9489        let mut weights = ReferenceWeights::new();
9490        let flat = |rows: &[[f32; 2]]| -> Vec<f32> { rows.iter().flatten().copied().collect() };
9491        weights.insert(
9492            layer_id(0, LayerTensor::KdaQuery),
9493            weight(&[2, 2], &flat(&wq)),
9494        );
9495        weights.insert(
9496            layer_id(0, LayerTensor::KdaKey),
9497            weight(&[2, 2], &flat(&wk)),
9498        );
9499        weights.insert(
9500            layer_id(0, LayerTensor::KdaValue),
9501            weight(&[2, 2], &flat(&wv)),
9502        );
9503        weights.insert(
9504            layer_id(0, LayerTensor::KdaQueryConv),
9505            weight(&[2, 2], &flat(&q_conv)),
9506        );
9507        weights.insert(
9508            layer_id(0, LayerTensor::KdaKeyConv),
9509            weight(&[2, 2], &flat(&k_conv)),
9510        );
9511        weights.insert(
9512            layer_id(0, LayerTensor::KdaValueConv),
9513            weight(&[2, 2], &flat(&v_conv)),
9514        );
9515        weights.insert(
9516            layer_id(0, LayerTensor::KdaForgetDown),
9517            weight(&[2, 2], &flat(&f_a)),
9518        );
9519        weights.insert(
9520            layer_id(0, LayerTensor::KdaForgetUp),
9521            weight(&[2, 2], &flat(&f_b)),
9522        );
9523        weights.insert(layer_id(0, LayerTensor::KdaDtBias), weight(&[2], &dt_bias));
9524        weights.insert(layer_id(0, LayerTensor::KdaALog), weight(&[1], &a_log));
9525        weights.insert(
9526            layer_id(0, LayerTensor::KdaBeta),
9527            weight(&[1, 2], &flat(&b_proj)),
9528        );
9529        weights.insert(
9530            layer_id(0, LayerTensor::KdaGateDown),
9531            weight(&[2, 2], &flat(&g_a)),
9532        );
9533        weights.insert(
9534            layer_id(0, LayerTensor::KdaGateUp),
9535            weight(&[2, 2], &flat(&g_b)),
9536        );
9537        weights.insert(
9538            layer_id(0, LayerTensor::KdaOutputNorm),
9539            weight(&[2], &o_norm),
9540        );
9541        weights.insert(
9542            layer_id(0, LayerTensor::KdaOutput),
9543            weight(&[2, 2], &flat(&wo)),
9544        );
9545
9546        let x_flat: Vec<f32> = x.iter().flatten().copied().collect();
9547        let (output, _) = kimi_delta_net(0, &plan, 1e-5, &weights, &x_flat, 3, 2).unwrap();
9548
9549        // Duplicated arithmetic, written independently of the operator.
9550        let sig = |value: f32| 1.0 / (1.0 + (-value).exp());
9551        let act = |value: f32| value * (1.0 / (1.0 + (-value).exp()));
9552        let mat2 = |m: &[[f32; 2]; 2], v: [f32; 2]| {
9553            [
9554                m[0][0] * v[0] + m[0][1] * v[1],
9555                m[1][0] * v[0] + m[1][1] * v[1],
9556            ]
9557        };
9558        let mut q_proj = [[0.0f32; 2]; 3];
9559        let mut k_proj = [[0.0f32; 2]; 3];
9560        let mut v_proj = [[0.0f32; 2]; 3];
9561        for token in 0..3 {
9562            q_proj[token] = mat2(&wq, x[token]);
9563            k_proj[token] = mat2(&wk, x[token]);
9564            v_proj[token] = mat2(&wv, x[token]);
9565        }
9566        let causal_conv = |proj: &[[f32; 2]; 3], conv: &[[f32; 2]; 2]| {
9567            let mut out = [[0.0f32; 2]; 3];
9568            for token in 0..3 {
9569                for channel in 0..2 {
9570                    let previous = if token == 0 {
9571                        0.0
9572                    } else {
9573                        proj[token - 1][channel]
9574                    };
9575                    out[token][channel] =
9576                        act(conv[channel][0] * previous + conv[channel][1] * proj[token][channel]);
9577                }
9578            }
9579            out
9580        };
9581        let mut q = causal_conv(&q_proj, &q_conv);
9582        let mut k = causal_conv(&k_proj, &k_conv);
9583        let v = causal_conv(&v_proj, &v_conv);
9584        for token in 0..3 {
9585            let q_inv = 1.0 / (q[token][0] * q[token][0] + q[token][1] * q[token][1] + 1e-6).sqrt();
9586            let k_inv = 1.0 / (k[token][0] * k[token][0] + k[token][1] * k[token][1] + 1e-6).sqrt();
9587            for channel in 0..2 {
9588                q[token][channel] *= q_inv * (1.0 / 2.0f32.sqrt());
9589                k[token][channel] *= k_inv;
9590            }
9591        }
9592        let decay_rate = a_log[0].exp();
9593        let mut expected = Vec::new();
9594        let mut state = [[0.0f32; 2]; 2];
9595        for token in 0..3 {
9596            let f_lin = mat2(&f_b, mat2(&f_a, x[token]));
9597            let g = [
9598                -5.0 * sig(decay_rate * (f_lin[0] + dt_bias[0])),
9599                -5.0 * sig(decay_rate * (f_lin[1] + dt_bias[1])),
9600            ];
9601            let beta = sig(b_proj[0][0] * x[token][0] + b_proj[0][1] * x[token][1]);
9602            for key_index in 0..2 {
9603                #[allow(clippy::needless_range_loop)]
9604                // allow: the explicit index loop keeps the offset arithmetic visible and aligned with the device-side indexing
9605                for value_index in 0..2 {
9606                    state[key_index][value_index] *= g[key_index].exp();
9607                }
9608            }
9609            let mut core = [0.0f32; 2];
9610            for value_index in 0..2 {
9611                let memory =
9612                    state[0][value_index] * k[token][0] + state[1][value_index] * k[token][1];
9613                let delta = (v[token][value_index] - memory) * beta;
9614                state[0][value_index] += k[token][0] * delta;
9615                state[1][value_index] += k[token][1] * delta;
9616            }
9617            for value_index in 0..2 {
9618                core[value_index] =
9619                    state[0][value_index] * q[token][0] + state[1][value_index] * q[token][1];
9620            }
9621            let gate = mat2(&g_b, mat2(&g_a, x[token]));
9622            let mean_square = (core[0] * core[0] + core[1] * core[1]) / 2.0;
9623            let inverse = 1.0 / (mean_square + 1e-5).sqrt();
9624            let gated = [
9625                core[0] * inverse * o_norm[0] * sig(gate[0]),
9626                core[1] * inverse * o_norm[1] * sig(gate[1]),
9627            ];
9628            let final_row = mat2(&wo, gated);
9629            expected.extend_from_slice(&final_row);
9630        }
9631        assert_eq!(output.len(), expected.len());
9632        for (index, (actual, wanted)) in output.iter().zip(&expected).enumerate() {
9633            assert!(
9634                (actual - wanted).abs() < 1e-5,
9635                "output[{index}] = {actual}, expected {wanted}"
9636            );
9637        }
9638    }
9639
9640    #[test]
9641    fn kpool_indexer_selects_causal_pools_and_appends_visible_tail() {
9642        use memra_gguf::model_plan::KpoolPlan;
9643
9644        let tokens = 8;
9645        let hidden = 2;
9646        let q_rank = 2;
9647        let identity = [1.0f32, 0.0, 0.0, 1.0];
9648        let mut weights = ReferenceWeights::new();
9649        weights.insert(
9650            layer_id(0, LayerTensor::SparseQuery),
9651            weight(&[2, 2], &identity),
9652        );
9653        weights.insert(
9654            layer_id(0, LayerTensor::SparseKey),
9655            weight(&[2, 2], &identity),
9656        );
9657        weights.insert(
9658            layer_id(0, LayerTensor::SparseKeyNorm),
9659            weight(&[2], &[1.0, 1.0]),
9660        );
9661        weights.insert(
9662            layer_id(0, LayerTensor::SparseKeyNormBias),
9663            weight(&[2], &[0.0, 0.0]),
9664        );
9665        weights.insert(
9666            layer_id(0, LayerTensor::SparseProjection),
9667            weight(&[1, 2], &[1.0, 1.0]),
9668        );
9669        weights.insert(
9670            layer_id(0, LayerTensor::SparseCompressorGate),
9671            weight(&[2, 2], &[0.3, -0.2, 0.1, 0.4]),
9672        );
9673        weights.insert(
9674            layer_id(0, LayerTensor::SparseCompressorPosition),
9675            weight(&[4, 2], &[0.1, 0.0, -0.1, 0.2, 0.05, -0.05, 0.0, 0.1]),
9676        );
9677        let x: Vec<f32> = (0..tokens * hidden)
9678            .map(|index| ((index % 5) as f32 - 2.0) * 0.3)
9679            .collect();
9680        let q_resid = x.clone();
9681
9682        // top_k 8 / pool 4 = a 2-pool budget, so every causally visible pool selects.
9683        let kpool = KpoolPlan {
9684            pool: 4,
9685            always_select_tail: true,
9686        };
9687        let allowed = kpool_allowed_tokens(
9688            0, 1, 2, 8, &kpool, &weights, &x, &q_resid, tokens, hidden, q_rank,
9689        )
9690        .unwrap();
9691        // Query 7 sees both complete pools; 8 % 4 == 0 leaves no tail.
9692        assert_eq!(allowed[7], (0..8).collect::<Vec<_>>());
9693        // Query 6: pool [4..=7] ends past the query, so only [0..=3] plus tail [4,5,6].
9694        assert_eq!(allowed[6], vec![0, 1, 2, 3, 4, 5, 6]);
9695        // Query 2 precedes any complete pool: tail only.
9696        assert_eq!(allowed[2], vec![0, 1, 2]);
9697
9698        // Without the tail, queries before the first complete pool have no candidates.
9699        let no_tail = KpoolPlan {
9700            pool: 4,
9701            always_select_tail: false,
9702        };
9703        let error = kpool_allowed_tokens(
9704            0, 1, 2, 8, &no_tail, &weights, &x, &q_resid, tokens, hidden, q_rank,
9705        )
9706        .unwrap_err();
9707        assert!(matches!(
9708            error,
9709            ReferenceError::InvalidPlan {
9710                reason: "k-pool selection produced an empty candidate set for a query",
9711                ..
9712            }
9713        ));
9714    }
9715}